#!/usr/bin/env python3
"""Independent algebra checks for the reviewed Hurwitz fibre theorems."""

from itertools import permutations, product


MATS = tuple(product(range(2), repeat=4))
IDENTITY = (1, 0, 0, 1)


def madd(a, b):
    return tuple(x + y for x, y in zip(a, b))


def mmul(a, b, modulus):
    return tuple(
        sum(a[2 * i + k] * b[2 * k + j] for k in range(2)) % modulus
        for i in range(2)
        for j in range(2)
    )


def lift(x, y=(0, 0, 0, 0)):
    return tuple((IDENTITY[i] + 2 * x[i] + 4 * y[i]) % 8 for i in range(4))


def two_layer_formula(xs, ys):
    linear = tuple(sum(x[k] for x in xs) for k in range(4))
    quadratic = tuple(sum(y[k] for y in ys) for k in range(4))
    for i in range(len(xs)):
        for j in range(i + 1, len(xs)):
            quadratic = madd(quadratic, mmul(xs[i], xs[j], 8))
    return tuple((IDENTITY[k] + 2 * linear[k] + 4 * quadratic[k]) % 8 for k in range(4))


def compose(p, q):
    return tuple(p[q[i]] for i in range(3))


def inverse(p):
    answer = [0] * 3
    for i, value in enumerate(p):
        answer[value] = i
    return tuple(answer)


S3 = tuple(permutations(range(3)))
NONIDENTITY = tuple(p for p in S3 if p != (0, 1, 2))


def hurwitz_pair(a, b):
    return b, compose(compose(inverse(b), a), b)


def hurwitz_pair_inverse(a, b):
    return compose(compose(a, b), inverse(a)), a


def tuple_product(values):
    answer = (0, 1, 2)
    for value in values:
        answer = compose(answer, value)
    return answer


def sign(p):
    inversions = sum(p[i] > p[j] for i in range(3) for j in range(i + 1, 3))
    return inversions % 2


def relative_nielsen_checks():
    checks = 0
    for factors in product(NONIDENTITY, repeat=3):
        target = tuple_product(factors)
        extended = factors + (inverse(target),)
        assert tuple_product(extended) == (0, 1, 2); checks += 1
        assert inverse(extended[-1]) == target; checks += 1
        for index in (0, 1):
            moved_pair = hurwitz_pair(factors[index], factors[index + 1])
            moved = factors[:index] + moved_pair + factors[index + 2:]
            assert tuple_product(moved) == target; checks += 1
            assert hurwitz_pair_inverse(*moved_pair) == factors[index:index + 2]; checks += 1
            assert tuple_product(moved + (inverse(target),)) == (0, 1, 2); checks += 1
            assert (sign(moved_pair[0]), sign(moved_pair[1])) == (sign(factors[index + 1]), sign(factors[index])); checks += 1
        colors = tuple(S3.index(factor) & 1 for factor in factors)
        receipt = sum(colors) % 2
        moved_colors = (colors[1], colors[0], colors[2])
        assert sum(moved_colors) % 2 == receipt; checks += 1
        conjugator = (1, 0, 2)
        conjugated = tuple(compose(compose(inverse(conjugator), factor), conjugator) for factor in factors)
        assert tuple_product(conjugated) == compose(compose(inverse(conjugator), target), conjugator); checks += 1
        assert sum(colors) % 2 == receipt; checks += 1
    assert checks == 1625
    return checks


def main():
    two_factor_checks = 0
    for x1, x2, y1, y2 in product(MATS, repeat=4):
        direct = mmul(lift(x1, y1), lift(x2, y2), 8)
        assert direct == two_layer_formula((x1, x2), (y1, y2))
        two_factor_checks += 1
    assert two_factor_checks == 65536

    three_factor_checks = 0
    zero = (0, 0, 0, 0)
    for x1, x2, x3 in product(MATS, repeat=3):
        direct = mmul(mmul(lift(x1), lift(x2), 8), lift(x3), 8)
        assert direct == two_layer_formula((x1, x2, x3), (zero, zero, zero))
        three_factor_checks += 1
    assert three_factor_checks == 4096

    e12 = (0, 1, 0, 0)
    e21 = (0, 0, 1, 0)
    assert mmul(e12, e21, 2) == (1, 0, 0, 0)
    assert mmul(e21, e12, 2) == (0, 0, 0, 1)
    assert two_layer_formula((e12, e21), (zero, zero)) != two_layer_formula((e21, e12), (zero, zero))

    nielsen_checks = relative_nielsen_checks()
    print("Hurwitz fibre audit: PASS")
    print(f"two-factor mod-eight identities: {two_factor_checks}")
    print(f"three-factor ordered-cross-term identities: {three_factor_checks}")
    print(f"total mod-eight identities: {two_factor_checks + three_factor_checks}")
    print("hostile order-sensitivity control: PASS")
    print(f"relative Nielsen and split-central identities: {nielsen_checks}")


if __name__ == "__main__":
    main()
