#!/usr/bin/env python3
"""Exact small-group cyclic-chain reference, not a spectrum/proof checker.

Enumerates every subgroup containing the supplied embedded subgroup, then
finds a shortest chain with normal inclusions and cyclic quotients. Limits
the input to multiplication tables of at most 16 elements and two million
table lookups. The new chromatic interpretation is only a source claim.
"""
import argparse
import json
from collections import deque
from pathlib import Path

MAX_ORDER = 16
MAX_LOOKUPS = 2_000_000
COMMIT = 'fd4aeeb2ee4fc729c18d98444fed42fd0529eeeb'


def calculate(record):
    if not isinstance(record, dict):
        raise ValueError('Input must be an object')
    prime = record.get('prime')
    if isinstance(prime, bool) or not isinstance(prime, int) or prime < 2 or prime > MAX_ORDER:
        raise ValueError('prime must be an integer between 2 and 16')
    if any(prime % d == 0 for d in range(2, int(prime ** 0.5) + 1)):
        raise ValueError('prime must be prime')
    table = record.get('multiplication_table')
    if not isinstance(table, list) or not 1 <= len(table) <= MAX_ORDER:
        raise ValueError('Use a multiplication table with 1 to 16 elements')
    n = len(table)
    if any(not isinstance(row, list) or len(row) != n for row in table):
        raise ValueError('Multiplication table must be square')
    if any(isinstance(x, bool) or not isinstance(x, int) or not 0 <= x < n for row in table for x in row):
        raise ValueError('Every product must be an integer element index')
    order_remainder = n
    while order_remainder > 1 and order_remainder % prime == 0:
        order_remainder //= prime
    if order_remainder != 1:
        raise ValueError('Group order must be a power of the supplied prime')
    lookups = 0
    def mul(a, b):
        nonlocal lookups
        lookups += 1
        if lookups > MAX_LOOKUPS:
            raise ValueError('Reference table-lookup budget exceeded')
        return table[a][b]
    identities = [e for e in range(n) if all(mul(e, x) == x and mul(x, e) == x for x in range(n))]
    if len(identities) != 1:
        raise ValueError('Table must have a unique two-sided identity')
    identity = identities[0]
    inverses = {}
    for x in range(n):
        candidates = [y for y in range(n) if mul(x, y) == identity and mul(y, x) == identity]
        if len(candidates) != 1:
            raise ValueError('Table must have unique two-sided inverses')
        inverses[x] = candidates[0]
    if any(mul(mul(a, b), c) != mul(a, mul(b, c)) for a in range(n) for b in range(n) for c in range(n)):
        raise ValueError('Multiplication table is not associative')
    supplied = record.get('subgroup')
    if not isinstance(supplied, list) or not supplied or any(isinstance(x, bool) or not isinstance(x, int) or not 0 <= x < n for x in supplied):
        raise ValueError('subgroup must be a nonempty list of element indices')
    if len(set(supplied)) != len(supplied):
        raise ValueError('Duplicate subgroup elements')
    h = frozenset(supplied)
    if identity not in h or any(inverses[x] not in h for x in h) or any(mul(a, b) not in h for a in h for b in h):
        raise ValueError('Supplied subset is not a subgroup')
    whole = frozenset(range(n))

    def generated(subgroup, element):
        generators = list(subgroup) + [element, inverses[element]]
        reached = {identity}
        queue = deque([identity])
        while queue:
            value = queue.popleft()
            for generator in generators:
                product = mul(value, generator)
                if product not in reached:
                    reached.add(product)
                    queue.append(product)
        return frozenset(reached)

    # Starting with H and adjoining one element recursively reaches all
    # overgroups of H; no enumeration of conjugacy classes replaces them.
    subgroups = {h}
    queue = deque([h])
    while queue:
        current = queue.popleft()
        for element in range(n):
            if element not in current:
                candidate = generated(current, element)
                if candidate not in subgroups:
                    subgroups.add(candidate)
                    queue.append(candidate)
    ordered = sorted(subgroups, key=lambda s: (len(s), tuple(sorted(s))))

    def cyclic_normal_step(a, b):
        if not a < b:
            return None
        if any(mul(mul(x, y), inverses[x]) not in a for x in b for y in a):
            return None
        index = len(b) // len(a)
        for generator in sorted(b - a):
            power = identity
            for exponent in range(1, index + 1):
                power = mul(power, generator)
                if power in a:
                    if exponent == index:
                        return dict(quotient_order=index, quotient_generator=generator)
                    break
        return None

    queue = deque([h])
    parent = {h: None}
    edges = {}
    while queue:
        current = queue.popleft()
        if current == whole:
            break
        for overgroup in ordered:
            if overgroup in parent:
                continue
            step = cyclic_normal_step(current, overgroup)
            if step is not None:
                parent[overgroup] = current
                edges[overgroup] = step
                queue.append(overgroup)
    if whole not in parent:
        raise AssertionError('A validated finite p-group should admit a cyclic subnormal chain')
    reverse_chain = [whole]
    while parent[reverse_chain[-1]] is not None:
        reverse_chain.append(parent[reverse_chain[-1]])
    chain = list(reversed(reverse_chain))
    return dict(group_order=n, prime=prime, identity=identity, subgroup=sorted(h),
                cyclic_length=len(chain)-1, chain=[sorted(s) for s in chain],
                quotient_steps=[edges[s] for s in chain[1:]],
                overgroups_enumerated=len(subgroups), table_lookups=lookups,
                arithmetic='exact_integer_table_operations', shortest_chain_method='exhaustive_overgroups_and_unweighted_BFS',
                limits=dict(maximum_order=MAX_ORDER, maximum_table_lookups=MAX_LOOKUPS),
                source_family='314', source_commit=COMMIT,
                interpretation='The source claims chromatic loss equals this cyclic length under finite p-local genuine-spectrum hypotheses; those hypotheses and the theorem are not verified by this calculator.',
                independent_proof_verification='not_run', spectrum_witness_constructed=False)


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('input', type=Path)
    parser.add_argument('--output', type=Path)
    args = parser.parse_args()
    result = json.dumps(calculate(json.loads(args.input.read_text())), indent=2, allow_nan=False) + '\n'
    if args.output:
        with args.output.open('x') as stream:
            stream.write(result)
    else:
        print(result, end='')


if __name__ == '__main__':
    main()
