"""Independent review of the starred-word shuffle quotient classifications.

The checker uses closed-form unary automata and direct labelled-map enumeration.
It does not import the original FPRD discovery scripts.  Only the Python
standard library is required.
"""

from __future__ import annotations

from dataclasses import dataclass
from itertools import product
import json


State = tuple
Edge = tuple[State, str, State]
ZERO: State = ("zero",)
EPSILON: State = ("epsilon",)


@dataclass(frozen=True)
class NFA:
    states: frozenset[State]
    edges: frozenset[Edge]
    initial: frozenset[State]
    final: frozenset[State]


def nxt(index: int, modulus: int) -> int:
    return index % modulus + 1


def location_automaton(m: int, n: int) -> NFA:
    rows = {i: ("r", i) for i in range(1, m + 1)}
    cols = {j: ("c", j) for j in range(1, n + 1)}
    grid = {(i, j): ("g", i, j) for i in range(1, m + 1) for j in range(1, n + 1)}
    edges: set[Edge] = {
        (ZERO, "a", rows[1]),
        (ZERO, "a", cols[1]),
    }
    for i in range(1, m + 1):
        edges.add((rows[i], "a", rows[nxt(i, m)]))
        edges.add((rows[i], "a", grid[i, 1]))
    for j in range(1, n + 1):
        edges.add((cols[j], "a", cols[nxt(j, n)]))
        edges.add((cols[j], "a", grid[1, j]))
    for i, j in product(range(1, m + 1), range(1, n + 1)):
        edges.add((grid[i, j], "a", grid[nxt(i, m), j]))
        edges.add((grid[i, j], "a", grid[i, nxt(j, n)]))
    return NFA(
        frozenset({ZERO, *rows.values(), *cols.values(), *grid.values()}),
        frozenset(edges),
        frozenset({ZERO}),
        frozenset({ZERO, rows[m], cols[n], grid[m, n]}),
    )


def prefix_automaton(m: int, n: int) -> NFA:
    core = {(p, r): ("q", p, r) for p in range(m) for r in range(n)}
    edges: set[Edge] = {(EPSILON, "a", core[0, 0])}
    for p, r in product(range(m), range(n)):
        edges.add((core[p, r], "a", core[(p + 1) % m, r]))
        edges.add((core[p, r], "a", core[p, (r + 1) % n]))
    return NFA(
        frozenset({EPSILON, *core.values()}),
        frozenset(edges),
        frozenset({EPSILON}),
        frozenset({EPSILON, core[m - 1, 0], core[0, n - 1]}),
    )


def published_projection(m: int, n: int) -> dict[State, State]:
    projection: dict[State, State] = {ZERO: EPSILON}
    for i in range(1, m + 1):
        projection[("r", i)] = ("q", i - 1, 0)
    for j in range(1, n + 1):
        projection[("c", j)] = ("q", 0, j - 1)
    for i, j in product(range(1, m + 1), range(1, n + 1)):
        projection[("g", i, j)] = ("q", i - 1, j % n)
    return projection


def quotient(source: NFA, projection: dict[State, State]) -> NFA:
    return NFA(
        frozenset(projection.values()),
        frozenset((projection[u], label, projection[v]) for u, label, v in source.edges),
        frozenset(projection[q] for q in source.initial),
        frozenset(projection[q] for q in source.final),
    )


def left_invariant(source: NFA, projection: dict[State, State]) -> bool:
    blocks: dict[State, list[State]] = {}
    for state, image in projection.items():
        blocks.setdefault(image, []).append(state)
    alphabet = {label for _, label, _ in source.edges}
    predecessors = {
        (state, label): {
            projection[u] for u, edge_label, v in source.edges
            if edge_label == label and v == state
        }
        for state in source.states
        for label in alphabet
    }
    for members in blocks.values():
        if len({state in source.initial for state in members}) != 1:
            return False
        for label in alphabet:
            if len({frozenset(predecessors[state, label]) for state in members}) != 1:
                return False
    return True


def direct_map_search(m: int, n: int) -> tuple[int, int, int]:
    source = location_automaton(m, n)
    target = prefix_automaton(m, n)
    source_core = sorted(source.states - {ZERO})
    target_core = sorted(target.states - {EPSILON})
    tested = ordinary_quotients = left_quotients = 0
    for images in product(target_core, repeat=len(source_core)):
        tested += 1
        projection = {ZERO: EPSILON, **dict(zip(source_core, images))}
        if set(projection.values()) != set(target.states):
            continue
        if quotient(source, projection) != target:
            continue
        ordinary_quotients += 1
        if left_invariant(source, projection):
            left_quotients += 1
    return tested, ordinary_quotients, left_quotients


def main() -> None:
    constructive_cases = 0
    left_certificate_cases = 0
    for m, n in product(range(1, 33), repeat=2):
        source = location_automaton(m, n)
        target = prefix_automaton(m, n)
        projection = published_projection(m, n)
        assert quotient(source, projection) == target
        constructive_cases += 1
        assert not left_invariant(source, projection)
    for m, n in product(range(1, 65), repeat=2):
        source = location_automaton(m, n)
        target = prefix_automaton(m, n)
        projection = published_projection(m, n)
        assert quotient(source, projection) == target
        if (m, n) == (1, 1):
            assert not left_invariant(source, projection)
        else:
            target_predecessors = {
                u for u, label, v in target.edges if label == "a" and v == ("q", 0, 0)
            }
            assert len(target_predecessors) == 3
        left_certificate_cases += 1

    alphabet = "abc"
    word_pairs = mixed_pairs = unary_controls = 0
    for total in range(2, 7):
        for m in range(1, total):
            n = total - m
            for x in map("".join, product(alphabet, repeat=m)):
                for y in map("".join, product(alphabet, repeat=n)):
                    word_pairs += 1
                    common_unary = len(set(x + y)) == 1
                    mixed_witness = any(a != b for a in x for b in y)
                    assert common_unary != mixed_witness
                    unary_controls += common_unary
                    mixed_pairs += mixed_witness

    boundary_cases = [(1, 1), (1, 2), (2, 1), (1, 3), (3, 1), (2, 2)]
    map_tests = ordinary_maps = left_maps = 0
    per_boundary = []
    for m, n in boundary_cases:
        tested, ordinary, left = direct_map_search(m, n)
        map_tests += tested
        ordinary_maps += ordinary
        left_maps += left
        per_boundary.append({"m": m, "n": n, "maps_tested": tested,
                             "ordinary_quotients": ordinary, "left_quotients": left})
    assert ordinary_maps > 0
    assert left_maps == 0

    print(json.dumps({
        "status": "pass",
        "constructive_unary_parameter_cases": constructive_cases,
        "left_obstruction_parameter_cases": left_certificate_cases,
        "general_word_pairs": word_pairs,
        "mixed_label_negative_witnesses": mixed_pairs,
        "common_letter_unary_controls": unary_controls,
        "direct_labelled_maps_tested": map_tests,
        "ordinary_quotient_maps_found": ordinary_maps,
        "left_quotient_maps_found": left_maps,
        "boundary_map_search": per_boundary,
    }, indent=2, sort_keys=True))


if __name__ == "__main__":
    main()
