"""Independent semantic review of the unary synchronization classifications.

This checker does not import the original FPRD audit.  It builds the strong,
arbitrary, and weak location/prefix automata directly from their transition
rules, compares their accepted unary lengths with a separate word-level event
semantics, and checks the published quotient maps and subset-orbit obstruction.
Only the Python standard library is required.
"""

from __future__ import annotations

from dataclasses import dataclass
from itertools import product
from math import gcd
import json


State = tuple
ZERO: State = ("zero",)
EPSILON: State = ("epsilon",)


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

    def step(self, subset: frozenset[State]) -> frozenset[State]:
        return frozenset(v for u, v in self.edges if u in subset)

    def accepts_length(self, length: int) -> bool:
        subset = self.initial
        for _ in range(length):
            subset = self.step(subset)
        return bool(subset & self.final)


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


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


def reachable_states(automaton: NFA) -> frozenset[State]:
    reached = automaton.initial
    frontier = automaton.initial
    while frontier:
        new = automaton.step(frontier) - reached
        reached |= new
        frontier = new
    return reached


def subset_orbit(automaton: NFA) -> tuple[int, int]:
    subset = automaton.initial
    seen: dict[frozenset[State], int] = {}
    while subset not in seen:
        seen[subset] = len(seen)
        subset = automaton.step(subset)
    return seen[subset], len(seen) - seen[subset]


def left_invariant(source: NFA, projection: dict[State, State]) -> bool:
    blocks: dict[State, set[State]] = {}
    for state, image in projection.items():
        blocks.setdefault(image, set()).add(state)
    predecessors = {
        q: {p for p, r in source.edges if r == q} for q in source.states
    }
    for members in blocks.values():
        if {q in source.initial for q in members} not in ({True}, {False}):
            return False
        profiles = {
            frozenset(projection[p] for p in predecessors[q]) for q in members
        }
        if len(profiles) != 1:
            return False
    return True


def strong_pair(m: int, n: int) -> tuple[NFA, NFA, dict[State, State]]:
    period = m * n // gcd(m, n)
    loc = [("loc", t % m, t % n) for t in range(period)]
    pre = [("pre", t % m, t % n) for t in range(period)]
    loc_edges = {(ZERO, loc[0]), *((loc[t], loc[(t + 1) % period]) for t in range(period))}
    pre_edges = {(EPSILON, pre[0]), *((pre[t], pre[(t + 1) % period]) for t in range(period))}
    source = NFA(
        frozenset({ZERO, *loc}),
        frozenset(loc_edges),
        frozenset({ZERO}),
        frozenset({ZERO, ("loc", m - 1, n - 1)}),
    )
    target = NFA(
        frozenset({EPSILON, *pre}),
        frozenset(pre_edges),
        frozenset({EPSILON}),
        frozenset({EPSILON, ("pre", m - 1, n - 1)}),
    )
    projection = {ZERO: EPSILON, **{q: ("pre", q[1], q[2]) for q in loc}}
    return source, target, projection


def arbitrary_location(m: int, n: int) -> NFA:
    left = {i: ("left", i) for i in range(m)}
    right = {j: ("right", j) for j in range(n)}
    pair = {(i, j): ("pair", i, j) for i in range(m) for j in range(n)}
    edges = {(ZERO, left[0]), (ZERO, right[0]), (ZERO, pair[0, 0])}
    for i in range(m):
        ip = nxt(i, m)
        edges |= {
            (left[i], left[ip]),
            (left[i], pair[i, 0]),
            (left[i], pair[ip, 0]),
        }
    for j in range(n):
        jp = nxt(j, n)
        edges |= {
            (right[j], right[jp]),
            (right[j], pair[0, j]),
            (right[j], pair[0, jp]),
        }
    for i, j in product(range(m), range(n)):
        ip, jp = nxt(i, m), nxt(j, n)
        edges |= {
            (pair[i, j], pair[ip, j]),
            (pair[i, j], pair[i, jp]),
            (pair[i, j], pair[ip, jp]),
        }
    return NFA(
        frozenset({ZERO, *left.values(), *right.values(), *pair.values()}),
        frozenset(edges),
        frozenset({ZERO}),
        frozenset({ZERO, left[m - 1], right[n - 1], pair[m - 1, n - 1]}),
    )


def arbitrary_prefix(m: int, n: int) -> NFA:
    q = {(i, j): ("q", i, j) for i in range(m) for j in range(n)}
    edges = {(EPSILON, q[0, 0])}
    for i, j in product(range(m), range(n)):
        ip, jp = nxt(i, m), nxt(j, n)
        edges |= {
            (q[i, j], q[ip, j]),
            (q[i, j], q[i, jp]),
            (q[i, j], q[ip, jp]),
        }
    return NFA(
        frozenset({EPSILON, *q.values()}),
        frozenset(edges),
        frozenset({EPSILON}),
        frozenset({EPSILON, q[m - 1, 0], q[0, n - 1], q[m - 1, n - 1]}),
    )


def arbitrary_projection(m: int, n: int) -> dict[State, State]:
    assert min(m, n) == 1
    projection: dict[State, State] = {ZERO: EPSILON}
    if m == 1:
        projection[("left", 0)] = ("q", 0, 0)
        for j in range(n):
            projection[("right", j)] = ("q", 0, j)
            projection[("pair", 0, j)] = ("q", 0, j)
    else:
        projection[("right", 0)] = ("q", 0, 0)
        for i in range(m):
            projection[("left", i)] = ("q", i, 0)
            projection[("pair", i, 0)] = ("q", i, 0)
    return projection


def weak_location(m: int, n: int) -> NFA:
    left = {i: ("left", i) for i in range(m)}
    right = {j: ("right", j) for j in range(n)}
    pair = {
        (mode, i, j): ("pair", mode, i, j)
        for mode in "NLR"
        for i in range(m)
        for j in range(n)
    }
    edges = {(ZERO, left[0]), (ZERO, right[0]), (ZERO, pair["N", 0, 0])}
    for i in range(m):
        edges |= {
            (left[i], left[nxt(i, m)]),
            (left[i], pair["N", nxt(i, m), 0]),
        }
    for j in range(n):
        edges |= {
            (right[j], right[nxt(j, n)]),
            (right[j], pair["N", 0, nxt(j, n)]),
        }
    for i, j in product(range(m), range(n)):
        ip, jp = nxt(i, m), nxt(j, n)
        edges |= {
            (pair["N", i, j], pair["N", ip, jp]),
            (pair["N", i, j], pair["L", ip, j]),
            (pair["N", i, j], pair["R", i, jp]),
            (pair["L", i, j], pair["N", ip, jp]),
            (pair["L", i, j], pair["L", ip, j]),
            (pair["R", i, j], pair["N", ip, jp]),
            (pair["R", i, j], pair["R", i, jp]),
        }
    return NFA(
        frozenset({ZERO, *left.values(), *right.values(), *pair.values()}),
        frozenset(edges),
        frozenset({ZERO}),
        frozenset(
            {
                ZERO,
                left[m - 1],
                right[n - 1],
                pair["N", m - 1, n - 1],
                pair["L", m - 1, n - 1],
                pair["R", m - 1, n - 1],
            }
        ),
    )


def weak_prefix(m: int, n: int) -> NFA:
    q = {
        (mode, i, j): ("pre", mode, i, j)
        for mode in "NLR"
        for i in range(m)
        for j in range(n)
    }
    edges = {(EPSILON, q[mode, 0, 0]) for mode in "NLR"}
    for i, j in product(range(m), range(n)):
        ip, jp = nxt(i, m), nxt(j, n)
        edges |= {(q["N", i, j], q[mode, ip, jp]) for mode in "NLR"}
        edges |= {(q["L", i, j], q[mode, ip, j]) for mode in "NL"}
        edges |= {(q["R", i, j], q[mode, i, jp]) for mode in "NR"}
    return NFA(
        frozenset({EPSILON, *q.values()}),
        frozenset(edges),
        frozenset({EPSILON}),
        frozenset(
            {
                EPSILON,
                q["N", m - 1, n - 1],
                q["L", m - 1, 0],
                q["R", 0, n - 1],
            }
        ),
    )


def weak_projection(m: int, n: int) -> dict[State, State]:
    projection: dict[State, State] = {ZERO: EPSILON}
    for i in range(m):
        projection[("left", i)] = ("pre", "L", i, 0)
    for j in range(n):
        projection[("right", j)] = ("pre", "R", 0, j)
    for i, j in product(range(m), range(n)):
        projection[("pair", "N", i, j)] = ("pre", "N", i, j)
        projection[("pair", "L", i, j)] = ("pre", "L", i, nxt(j, n))
        projection[("pair", "R", i, j)] = ("pre", "R", nxt(i, m), j)
    return projection


def word_semantics_accepts(kind: str, m: int, n: int, length: int) -> bool:
    states = {(0, 0, "N")}
    for _ in range(length):
        following = set()
        for left, right, mode in states:
            if kind == "strong":
                following.add((left + 1, right + 1, "N"))
            elif kind == "arbitrary":
                following |= {
                    (left + 1, right, "N"),
                    (left, right + 1, "N"),
                    (left + 1, right + 1, "N"),
                }
            elif kind == "weak":
                following.add((left + 1, right + 1, "N"))
                if mode in "NL":
                    following.add((left + 1, right, "L"))
                if mode in "NR":
                    following.add((left, right + 1, "R"))
            else:
                raise ValueError(kind)
        states = following
    return any(left % m == 0 and right % n == 0 for left, right, _ in states)


def run() -> dict:
    strong_cases = 0
    arbitrary_cases = 0
    weak_cases = 0
    for m, n in product(range(1, 25), repeat=2):
        source, target, projection = strong_pair(m, n)
        assert quotient(source, projection) == target
        assert left_invariant(source, projection)
        assert reachable_states(source) == source.states
        strong_cases += 1

        source, target = arbitrary_location(m, n), arbitrary_prefix(m, n)
        assert reachable_states(source) == source.states
        if min(m, n) == 1:
            projection = arbitrary_projection(m, n)
            assert quotient(source, projection) == target
            assert left_invariant(source, projection) == (m == n == 1)
        else:
            q00 = ("q", 0, 0)
            assert (q00, q00) not in target.edges
            assert (("left", 0), ("pair", 0, 0)) in source.edges
        arbitrary_cases += 1

    orbit_cases = 0
    for m, n in product(range(1, 33), repeat=2):
        source, target = weak_location(m, n), weak_prefix(m, n)
        projection = weak_projection(m, n)
        assert quotient(source, projection) == target
        assert reachable_states(source) == source.states
        weak_cases += 1
        if m <= 16 and n <= 16:
            assert subset_orbit(source) == (m + n, m * n // gcd(m, n))
            expected = (1, 1) if (m, n) == (1, 1) else (m + n, 1)
            assert subset_orbit(target) == expected
            orbit_cases += 1

    semantic_comparisons = 0
    for m, n in product(range(1, 7), repeat=2):
        automata = {
            "strong": strong_pair(m, n)[:2],
            "arbitrary": (arbitrary_location(m, n), arbitrary_prefix(m, n)),
            "weak": (weak_location(m, n), weak_prefix(m, n)),
        }
        for length in range(0, 2 * (m + n) + 7):
            for kind, pair in automata.items():
                expected = word_semantics_accepts(kind, m, n, length)
                assert all(a.accepts_length(length) == expected for a in pair)
                semantic_comparisons += 2

    return {
        "status": "pass",
        "strong": {
            "parameter_cases": strong_cases,
            "result": "accessible location and prefix automata are isomorphic",
        },
        "arbitrary": {
            "parameter_cases": arbitrary_cases,
            "ordinary_quotient_iff": "min(m,n)=1",
            "left_quotient_iff": "m=n=1",
            "negative_control": "the forced initial block creates a forbidden self-loop at (2,2)",
        },
        "weak": {
            "quotient_parameter_cases": weak_cases,
            "subset_orbit_cases": orbit_cases,
            "ordinary_quotient": "all parameter cases",
            "left_quotient": "no parameter case",
            "exceptional_control": "at (1,1), periods agree but the location orbit has one extra transient subset",
        },
        "word_semantics": {
            "automaton_acceptance_comparisons": semantic_comparisons,
            "parameter_box": "1 <= m,n <= 6",
            "length_bound": "0 <= k < 2(m+n)+7",
            "result": "both automata agree with an independent event-sequence oracle",
        },
    }


if __name__ == "__main__":
    print(json.dumps(run(), indent=2, sort_keys=True))
