"""Checks and a report for Part 5: the floor with a twist.

Run from this folder: python test_mobius.py
"""
import random
import statistics
import time
from collections import deque

import blind_torus
import blind_twist
import mobius
from parity_dsu import ParityDSU
from surface import Robot, flat, frames_reachable, klein, simulate, torus
from surface import mobius as mobius_surface

STRIP = ['.......', '.......', '...*...', '.......', '.......']      # 5 rows, 7 columns
BLOCKED = ['.......', '.......', '*..#...', '.......', '.......']    # a pillar on the way


def random_floor(rng, R, C, furniture):
    grid = [['#' if rng.random() < furniture else '.' for _ in range(C)] for _ in range(R)]
    r, c = rng.randrange(R), rng.randrange(C)
    grid[r][c] = '*'
    return [''.join(row) for row in grid]


def reach(s, start):
    return {f for f, _ in frames_reachable(s, start)}


def test_walk(rng):
    """The walk in the robot's own directions cleans everything; the same walk in the map's
    directions goes wrong as soon as it crosses the seam and then moves up or down."""
    wrong = crossed = 0
    for _ in range(500):
        room = random_floor(rng, rng.randint(1, 9), rng.randint(3, 12), rng.random() * 0.4)
        s, start = mobius_surface(room)
        cmds = mobius.solution(room)
        cleaned, _, _ = simulate(s, start, cmds)
        assert cleaned == reach(s, start) and len(cmds) == 2 * (len(cleaned) - 1)
        naive = mobius.map_walk(s, start)
        mistake = mobius.first_mistake(s, start, naive)
        if mistake is not None:
            wrong += 1
        if any(g and g[2] for f in reach(s, start) for g in s.glue[f]):
            crossed += 1
        else:
            assert mistake is None                  # no seam to cross, nothing goes wrong
    s, start = mobius_surface(STRIP)
    _, field, frame = simulate(s, start, '>' * 7)
    assert field == start and frame == (2, True)    # home, upside down
    _, field, _ = simulate(s, start, '^' + '>' * 7)
    assert s.names[field] == (3, 3)                 # one row up, then round: one row down
    print(f'walk .................... ok (500 rooms; the map\'s directions went wrong on {wrong}, '
          f'of the {crossed} whose floor reaches the seam)')


def test_orientable(rng):
    """Two frames per field exactly when some loop mirrors the robot, and the union-find
    agrees with that search on every room."""
    counts = {True: 0, False: 0}
    for _ in range(600):
        build = rng.choice([mobius_surface, klein, torus, flat])
        room = random_floor(rng, rng.randint(3, 9), rng.randint(3, 10), rng.random() * 0.55)
        s, start = build(room)
        o = mobius.orientable(s, start)
        counts[o] += 1
        dsu = ParityDSU()
        for f in reach(s, start):
            dsu.add(f)
        for f in reach(s, start):
            for g in s.glue[f]:
                if g is not None:
                    dsu.union(f, g[0], int(g[2]))
        assert dsu.consistent == o
        pairs = frames_reachable(s, start)
        assert len({(f, m) for f, (_, m) in pairs}) == len(reach(s, start)) * (1 if o else 2)
    print(f'orientable .............. ok (600 rooms: {counts[True]} orientable, '
          f'{counts[False]} with a loop that mirrors)')


def test_dsu(rng):
    """Union-find against a fresh search after every union, on random graphs with random
    claims."""
    for _ in range(300):
        n = rng.randint(2, 30)
        dsu, edges = ParityDSU(), []
        for x in range(n):
            dsu.add(x)
        for _ in range(rng.randint(1, 60)):
            a, b, p = rng.randrange(n), rng.randrange(n), rng.randrange(2)
            dsu.union(a, b, p)
            edges.append((a, b, p))
            assert dsu.consistent == consistent(n, edges)   # once contradicted, always
            for x in range(n):
                for y in range(n):
                    got = dsu.same_frame(x, y)
                    want = parity_between(n, edges, x, y)
                    if dsu.consistent:
                        assert got == want
    print('union-find .............. ok (300 random graphs, checked against a search each step)')


def consistent(n, edges):
    return all(parity_between(n, edges, a, b) in (None, p == 0) for a, b, p in edges)


def parity_between(n, edges, a, b):
    """By search: None if unconnected, else whether some path from a to b has even parity.
    (With no odd loop, every path has the same parity.)"""
    adj = [[] for _ in range(n)]
    for x, y, p in edges:
        adj[x].append((y, p))
        adj[y].append((x, p))
    seen = {(a, 0)}
    queue = deque([(a, 0)])
    while queue:
        x, q = queue.popleft()
        for y, p in adj[x]:
            if (y, q ^ p) not in seen:
                seen.add((y, q ^ p))
                queue.append((y, q ^ p))
    if (b, 0) in seen:
        return True
    return False if (b, 1) in seen else None


def first_twist_dsu(s, order):
    """Floor appears field by field in `order`. Returns how many fields were there when the
    first loop that mirrors the robot closed, and how many parent links find followed."""
    dsu, present = ParityDSU(), set()
    for k, f in enumerate(order, 1):
        dsu.add(f)
        present.add(f)
        for g in s.glue[f]:
            if g is not None and g[0] in present and dsu.union(f, g[0], int(g[2])) == 'contradicts':
                return k
    return None


def first_twist_search(s, order):
    """The same, the slow way: after each new field, search the pieces it touches for two
    ways round the same field."""
    present = set()
    for k, f in enumerate(order, 1):
        present.add(f)
        seen = {(f, 0)}
        todo = [(f, 0)]
        while todo:
            x, m = todo.pop()
            for g in s.glue[x]:
                if g is not None and g[0] in present:
                    y = (g[0], m ^ int(g[2]))
                    if (g[0], 1 - y[1]) in seen:
                        return k
                    if y not in seen:
                        seen.add(y)
                        todo.append(y)
    return None


def report_twist(rng):
    print('\nFloor added one field at a time, in random order, on an n x n Mobius strip:')
    print('  n     fields present when the first mirroring loop closed (mean, as a share)')
    for n in (8, 16, 32, 64, 128):
        s, _ = mobius_surface(['.' * n] * n)
        shares = []
        for _ in range(400 if n <= 64 else 100):
            order = list(range(n * n))
            rng.shuffle(order)
            shares.append(first_twist_dsu(s, order) / (n * n))
        print(f'  {n:<5} {statistics.mean(shares):.3f} (median {statistics.median(shares):.3f})')
    s, _ = mobius_surface(['.' * 32] * 32)
    for _ in range(20):
        order = list(range(32 * 32))
        rng.shuffle(order)
        assert first_twist_dsu(s, order) == first_twist_search(s, order)
    order = list(range(64 * 64))
    rng.shuffle(order)
    s, _ = mobius_surface(['.' * 64] * 64)
    t0 = time.perf_counter()
    a = first_twist_dsu(s, order)
    t1 = time.perf_counter()
    b = first_twist_search(s, order)
    t2 = time.perf_counter()
    assert a == b
    print(f'  one 64 x 64 strip: union-find {1000 * (t1 - t0):.0f} ms, search after every field '
          f'{1000 * (t2 - t1):.0f} ms (same answer, {a} fields)')


def test_blind(rng):
    """With the arrow on the dock, the robot cleans every room and stops; with a round dock it
    folds its map wrongly on a Mobius strip."""
    by = {}
    for build in (mobius_surface, klein, torus, flat):
        for _ in range(120):
            room = random_floor(rng, rng.randint(3, 8), rng.randint(3, 9), rng.random() * 0.4)
            s, start = build(room)
            everything = reach(s, start)
            frame = (rng.randrange(4), rng.random() < 0.5)
            robot = Robot(s, start, frame, limit=10_000_000)
            count, stats = blind_twist.solution(robot)
            assert count == len(everything) and robot.cleaned == everything
            assert stats['group'].mirrored() == (not mobius.orientable(s, start))
            by.setdefault(build.__name__, []).append(robot.moves / max(1, len(everything) - 1))
    print('blind robot ............. ok (' + ', '.join(
        f'{k} {statistics.mean(v):.2f}' for k, v in by.items()) + ' moves per field)')


def round_dock(rng, tries=150):
    """Part 4's robot, which only knows whether it is on the dock, on Mobius strips."""
    outcomes = {'right': 0, 'wrong count': 0, 'map contradiction': 0, 'never stopped': 0}
    for _ in range(tries):
        room = random_floor(rng, rng.randint(3, 8), rng.randint(3, 9), rng.random() * 0.4)
        s, start = mobius_surface(room)
        robot = Robot(s, start, limit=200_000)

        class RoundDock:
            move = staticmethod(robot.move)

            @staticmethod
            def on_dock():
                return robot.on_dock() is not None
        try:
            count, _ = blind_torus.solution(RoundDock)
            outcome = 'right' if count == len(reach(s, start)) else 'wrong count'
        except AssertionError:
            outcome = 'map contradiction'
        except RuntimeError:
            outcome = 'never stopped'
        outcomes[outcome] += 1
    print(f'\nPart 4\'s robot with a round dock, {tries} Mobius rooms: {outcomes}')


if __name__ == '__main__':
    rng = random.Random(5)
    test_walk(rng)
    test_orientable(rng)
    test_dsu(rng)
    test_blind(rng)
    round_dock(rng)
    report_twist(rng)
