#!/usr/bin/env python3
"""Dependency-free audit of the 2P3 type-slot synchronization obstruction."""

from collections import Counter
from itertools import combinations, permutations, product


def relabel_edges(edges, order):
    return frozenset(
        tuple(sorted((order.index(a), order.index(b)))) for a, b in edges
    )


def canonical_graph(vertices, edges):
    vertices = tuple(vertices)
    n = len(vertices)
    best = None
    for order in permutations(vertices):
        bits = tuple(
            int(tuple(sorted((order[i], order[j]))) in edges)
            for i in range(n)
            for j in range(i + 1, n)
        )
        if best is None or bits < best:
            best = bits
    return best


def unlabeled_graphs(n):
    pairs = list(combinations(range(n), 2))
    representatives = {}
    for mask in range(1 << len(pairs)):
        edges = frozenset(pairs[k] for k in range(len(pairs)) if mask >> k & 1)
        key = canonical_graph(range(n), edges)
        representatives.setdefault(key, edges)
    return list(representatives.values())


def deletion_partition(n, edges):
    classes = {}
    for vertex in range(n):
        remaining = [v for v in range(n) if v != vertex]
        classes.setdefault(canonical_graph(remaining, edges), set()).add(vertex)
    return {frozenset(block) for block in classes.values()}


def automorphism_partition(n, edges):
    parent = list(range(n))

    def find(x):
        while parent[x] != x:
            parent[x] = parent[parent[x]]
            x = parent[x]
        return x

    def union(a, b):
        a, b = find(a), find(b)
        if a != b:
            parent[b] = a

    for order in permutations(range(n)):
        image = frozenset(tuple(sorted((order[a], order[b]))) for a, b in edges)
        if image == edges:
            for a, b in enumerate(order):
                union(a, b)
    blocks = {}
    for v in range(n):
        blocks.setdefault(find(v), set()).add(v)
    return {frozenset(block) for block in blocks.values()}


VERTICES = tuple(range(6))
EDGES = frozenset({(0, 1), (1, 2), (3, 4), (4, 5)})
CENTERS = (1, 4)
LEAVES = (0, 2, 3, 5)
SLOTS = tuple(combinations(VERTICES, 2))


def double_deletion_type(i, j):
    return canonical_graph([v for v in VERTICES if v not in (i, j)], EDGES)


ROW_PROFILES = {
    i: Counter(double_deletion_type(i, j) for j in VERTICES if j != i)
    for i in VERTICES
}
COLORS = tuple(sorted({double_deletion_type(i, j) for i, j in SLOTS}))


def synchronizations():
    remaining = {i: ROW_PROFILES[i].copy() for i in VERTICES}
    answers = []

    def search(k, current):
        if k == len(SLOTS):
            answers.append(dict(current))
            return
        i, j = SLOTS[k]
        for color in COLORS:
            if remaining[i][color] and remaining[j][color]:
                remaining[i][color] -= 1
                remaining[j][color] -= 1
                current[i, j] = color
                search(k + 1, current)
                del current[i, j]
                remaining[i][color] += 1
                remaining[j][color] += 1

    search(0, {})
    return answers


def row_namings(sync, omitted):
    labels = [v for v in VERTICES if v != omitted]
    answers = []
    for assigned_vertices in permutations(labels):
        naming = dict(zip(labels, assigned_vertices))
        if all(
            canonical_graph(
                [v for v in labels if v != naming[slot]],
                EDGES,
            )
            == sync[tuple(sorted((omitted, slot)))]
            for slot in labels
        ):
            answers.append(naming)
    return answers


def coherent_naming_count(sync):
    options = [row_namings(sync, omitted) for omitted in VERTICES]
    assert [len(row) for row in options] == [4] * 6
    coherent = 0
    for chosen in product(*options):
        seen = {}
        valid = True
        for omitted, naming in enumerate(chosen):
            labels = [v for v in VERTICES if v != omitted]
            for a, b in combinations(labels, 2):
                value = tuple(sorted((naming[a], naming[b]))) in EDGES
                if (a, b) in seen and seen[a, b] != value:
                    valid = False
                    break
                seen[a, b] = value
            if not valid:
                break
        coherent += int(valid)
    return coherent


def matching_pair(sync):
    # P is the matching formed by the two leaf-leaf slots of path-plus-isolate type.
    path_plus_isolate = canonical_graph((0, 1, 2, 3), {(0, 1), (1, 2)})
    p = frozenset(
        pair for pair in combinations(LEAVES, 2) if sync[pair] == path_plus_isolate
    )
    # Q pairs leaves assigned to the same center by the one-edge slot type.
    one_edge = canonical_graph((0, 1, 2, 3), {(0, 1)})
    center_blocks = []
    for center in CENTERS:
        block = tuple(sorted(v for v in LEAVES if sync[tuple(sorted((center, v)))] == one_edge))
        assert len(block) == 2
        center_blocks.append(block)
    q = frozenset(center_blocks)
    assert len(p) == len(q) == 2
    return p, q


MATCHINGS = tuple(
    frozenset((tuple(sorted((LEAVES[0], LEAVES[i]))), tuple(sorted(set(LEAVES) - {LEAVES[0], LEAVES[i]}))))
    for i in range(1, 4)
)
MATCHING_VECTOR = {matching: vector for matching, vector in zip(MATCHINGS, ((1, 0), (0, 1), (1, 1)))}


def xor(a, b):
    return (a[0] ^ b[0], a[1] ^ b[1])


def defect(sync):
    p, q = matching_pair(sync)
    return xor(MATCHING_VECTOR[p], MATCHING_VECTOR[q])


def act_on_sync(sync, permutation):
    return {
        tuple(sorted((permutation[i], permutation[j]))): color
        for (i, j), color in sync.items()
    }


def matching_action(permutation, matching):
    return frozenset(tuple(sorted((permutation[a], permutation[b]))) for a, b in matching)


def main():
    syncs = synchronizations()
    assert len(syncs) == 18
    keys = {tuple(sync[slot] for slot in SLOTS): i for i, sync in enumerate(syncs)}

    symmetries = []
    for center_order in permutations(CENTERS):
        for leaf_order in permutations(LEAVES):
            permutation = dict(zip(CENTERS, center_order)) | dict(zip(LEAVES, leaf_order))
            symmetries.append(permutation)
    assert len(symmetries) == 48

    unseen = set(range(len(syncs)))
    orbit_sizes = []
    while unseen:
        seed = next(iter(unseen))
        orbit = {
            keys[tuple(act_on_sync(syncs[seed], symmetry)[slot] for slot in SLOTS)]
            for symmetry in symmetries
        }
        orbit_sizes.append(len(orbit))
        unseen -= orbit
    assert sorted(orbit_sizes) == [6, 12]

    coherence = [coherent_naming_count(sync) for sync in syncs]
    assert Counter(coherence) == Counter({4096: 6, 0: 12})
    assert all((count > 0) == (matching_pair(sync)[0] == matching_pair(sync)[1]) for sync, count in zip(syncs, coherence))

    defects = Counter(defect(sync) for sync in syncs)
    assert defects == Counter({(0, 0): 6, (1, 0): 4, (0, 1): 4, (1, 1): 4})

    equivariance_checks = 0
    for sync in syncs:
        p, q = matching_pair(sync)
        for symmetry in symmetries:
            moved = act_on_sync(sync, symmetry)
            moved_p = matching_action(symmetry, p)
            moved_q = matching_action(symmetry, q)
            assert matching_pair(moved) == (moved_p, moved_q)
            assert defect(moved) == xor(MATCHING_VECTOR[moved_p], MATCHING_VECTOR[moved_q])
            equivariance_checks += 1
    assert equivariance_checks == 864

    small_graphs = 0
    for n in range(1, 6):
        graphs = unlabeled_graphs(n)
        small_graphs += len(graphs)
        for edges in graphs:
            assert deletion_partition(n, edges) == automorphism_partition(n, edges)
    assert small_graphs == 52

    print("graph-deck audit: PASS")
    print(f"type-slot synchronizations: {len(syncs)}")
    print(f"row-profile symmetry orbits: {sorted(orbit_sizes)}")
    print(f"candidate row namings per synchronization: {4 ** 6}")
    print(f"coherent naming distribution: {dict(sorted(Counter(coherence).items()))}")
    print(f"matching-defect distribution: {dict(sorted(defects.items()))}")
    print(f"full S2 x S4 equivariance checks: {equivariance_checks}")
    print(f"unlabeled marked-orbit controls through order five: {small_graphs}")


if __name__ == "__main__":
    main()
