"""Checks and a report for Part 4: the room that wraps around.

Run from this folder: python test_torus.py
"""
import math
import random
import statistics
from itertools import combinations
from math import gcd
from types import SimpleNamespace

import blind_robot
import blind_torus
import periods
import route_exact
import torus
from surface import simulate
from surface import torus as torus_surface
from two_colour import is_odd_loop, two_colour

# Rooms from the post.
EMPTY_3 = ['...', '*..', '...']                     # the checkerboard bound says 9; 8 is enough
NOTCH = ['.*.', '#.#', '...']                       # Part 3's pruning finds 7; 6 is enough
CORRIDOR = ['##.##', '##.##', '*..##', '##.##', '##.##']    # an endless corridor, to the robot
POCKET = ['#####', '#*..#', '#.#.#', '#...#', '#####']      # no loop goes round: rank 0


def test_walk(rng):
    """Part 1's walk on a torus cleans everything in 2(N - 1) commands."""
    for _ in range(500):
        room = torus.random_room(rng, rng.randint(3, 15), rng.randint(3, 15), rng.random() * 0.6)
        s, start = torus_surface(room)
        cmds = torus.solution(room)
        cleaned, _, _ = simulate(s, start, cmds)
        assert {s.names[f] for f in cleaned} == torus.floor(room)
        assert len(cmds) == 2 * (len(cleaned) - 1)
    print('walk .................... ok (500 rooms)')


def test_hermite(rng):
    """The normal form is canonical and spans the same lattice as its input."""
    for _ in range(3000):
        vs = [(rng.randint(-30, 30), rng.randint(-30, 30)) for _ in range(rng.randint(1, 5))]
        basis = periods.hermite(vs)
        for v in vs:                                # every input is in the lattice...
            assert periods.reduce(v, basis) == (0, 0)
        dets = [abs(a[0] * b[1] - a[1] * b[0]) for a, b in combinations(vs, 2)]
        g = 0
        for d in dets:
            g = gcd(g, d)
        if len(basis) == 2:                         # ...and the area of a cell is right
            assert basis[0][0] * basis[1][1] == g and 0 <= basis[0][1] < basis[1][1]
        rng.shuffle(vs)
        assert periods.hermite(vs) == basis         # the order of the input doesn't matter
        extra = [(u[0] + 3 * w[0], u[1] + 3 * w[1]) for u, w in zip(vs, vs[1:])]
        assert periods.hermite(vs + extra) == basis  # nor does adding redundant periods
    assert periods.periods(['*' + '.' * 6] + ['.' * 7] * 4) == ((5, 0), (0, 7))
    assert periods.periods(CORRIDOR) == ((5, 0),)
    assert periods.periods(POCKET) == ()
    print('hermite normal form ..... ok (3,000 random lattices)')


def part3_stops(room, limit=20_000):
    robot = torus.Robot(room, limit)
    try:
        return blind_robot.solution(SimpleNamespace(move=robot.move)), robot
    except RuntimeError:
        return None, robot


def test_prediction(rng):
    """Part 3's blind robot stops exactly on the rooms whose periods have rank 0."""
    ranks = [0, 0, 0]
    for _ in range(600):
        room = torus.random_room(rng, rng.randint(3, 12), rng.randint(3, 12), rng.random() * 0.6)
        rank = len(periods.periods(room))
        ranks[rank] += 1
        count, _ = part3_stops(room)
        assert (count is not None) == (rank == 0), room
        if count is not None:
            assert count == len(torus.floor(room))
    assert part3_stops(CORRIDOR)[0] is None and part3_stops(POCKET)[0] == 8
    print(f'prediction .............. ok (600 rooms, rank 0/1/2: {ranks[0]}/{ranks[1]}/{ranks[2]})')
    return ranks


def cover(room, basis):
    """The room twice over, stacked in a direction some loop goes round an odd number of times,
    so the doubled floor is still in one piece. The second copy's dock is left out. Returns
    None if every loop goes round an even number of times both ways."""
    R, C = len(room), len(room[0])
    down = gcd(*(a for a, _ in basis)) if basis else 0      # vertical parts of all periods
    across = gcd(*(b for _, b in basis)) if basis else 0    # horizontal parts
    if down and (down // R) % 2:
        return room + [row.replace('*', '.') for row in room]
    if across and (across // C) % 2:
        return [row + row.replace('*', '.') for row in room]
    return None


def test_covers(rng):
    """On a room with a loop round the torus and on its double, Part 3's robot does exactly the
    same thing, move for move, for as long as anyone cares to watch."""
    checked = 0
    for _ in range(300):
        room = torus.random_room(rng, rng.randint(3, 10), rng.randint(3, 10), rng.random() * 0.5)
        basis = periods.periods(room)
        if not basis:
            continue
        big = cover(room, basis)
        if big is None:
            continue
        assert len(torus.floor(big)) == 2 * len(torus.floor(room))
        _, a = part3_stops(room, 20_000)
        _, b = part3_stops(big, 20_000)
        assert a.log == b.log
        checked += 1
    print(f'covers .................. ok ({checked} rooms, identical for 20,000 move calls)')


def run_landmark(room, limit=10_000_000):
    robot = torus.Robot(room, limit)
    lattice = blind_torus.Lattice()
    count, stats = blind_torus.solution(robot, lattice)
    assert count == len(torus.floor(room)) and robot.cleaned == torus.floor(room)
    assert lattice.basis == periods.periods(room)   # it learned the room's periods exactly
    return robot, stats


def test_landmark(rng):
    """With a dock to recognise, the robot cleans every room and stops, and learns the room's
    lattice of periods exactly."""
    rows = []
    for _ in range(400):
        room = torus.random_room(rng, rng.randint(3, 14), rng.randint(3, 14), rng.random() * 0.55)
        n = len(torus.floor(room))
        robot, stats = run_landmark(room)
        rows.append((len(periods.periods(room)), n, robot, stats))
        basis = periods.periods(room)
        big = cover(room, basis) if basis else None
        if big is not None:                         # its double is a different room: it knows
            robot2, _ = run_landmark(big)
            assert len(robot2.cleaned) == 2 * n
    robot, stats = run_landmark(CORRIDOR)
    assert stats['periods'] == 2                    # (10, 0) first, then (5, 0)
    print(f'landmark robot .......... ok ({len(rows)} rooms and their doubles)')
    return rows


def test_two_colour(rng):
    """A colouring or a checked odd loop, and the periods predict which: a loop's length has
    the parity of its period's two entries added."""
    counts = {'colours': 0, 'odd loop': 0}
    for _ in range(600):
        room = torus.random_room(rng, rng.randint(3, 11), rng.randint(3, 11), rng.random() * 0.6)
        s, start = torus_surface(room)
        kind, result = two_colour(s, start)
        counts[kind] += 1
        predicted = all((a + b) % 2 == 0 for a, b in periods.periods(room))
        assert (kind == 'colours') == predicted, room
        if kind == 'odd loop':
            assert is_odd_loop(s, result)
        else:
            assert all(result[f] != result[g[0]] for f in result for g in s.glue[f] if g)
    print(f'two colours ............. ok (600 rooms: {counts["colours"]} coloured, '
          f'{counts["odd loop"]} odd loops checked)')


def test_exact(rng):
    """With a proper colouring, or none, the pruned search agrees with brute force. Trusting
    (row + column) mod 2 on an odd torus does not."""
    wrong_prune = wrong_bound = rooms = 0
    for _ in range(1500):
        room = torus.random_room(rng, rng.choice([3, 4, 5]), rng.randint(3, 6), rng.random() * 0.5)
        s, start = torus_surface(room)
        n = len(route_exact.floor_from(s, start))
        if not 3 <= n <= 13:
            continue
        rooms += 1
        referee = route_exact.fewest(s, start, prune=False)
        kind, colours = two_colour(s, start)
        right = route_exact.fewest(s, start, colour=colours if kind == 'colours' else None)
        assert len(right) == len(referee)
        assert simulate(s, start, right)[0] == set(route_exact.floor_from(s, start))
        naive = {f: (r + c) % 2 for f, (r, c) in enumerate(s.names)}
        wrong_prune += len(route_exact.fewest(s, start, colour=naive)) > len(referee)
        wrong_bound += route_exact.lower_bound(s, start, naive) > len(referee)
    for room, best, bound, pruned in ((EMPTY_3, 8, 9, 8), (NOTCH, 6, 9, 7)):
        s, start = torus_surface(room)
        naive = {f: (r + c) % 2 for f, (r, c) in enumerate(s.names)}
        assert len(route_exact.fewest(s, start, prune=False)) == best
        assert route_exact.lower_bound(s, start, naive) == bound
        assert len(route_exact.fewest(s, start, colour=naive)) == pruned
    print(f'fewest commands ......... ok ({rooms} small rooms: Part 3\'s colouring overstated the '
          f'bound on {wrong_bound} and pruned the best route away on {wrong_prune})')
    return rooms, wrong_bound, wrong_prune


def report(rows):
    print('\nThe landmark robot, by the rank of the room\'s periods (moves per field, N - 1):')
    for rank in (0, 1, 2):
        mine = [r for r in rows if r[0] == rank]
        m = [robot.moves / (n - 1) for _, n, robot, _ in mine if n > 1]
        print(f'  rank {rank}: {len(mine):>3} rooms, moves/(N-1) mean {statistics.mean(m):.2f}, '
              f'max {max(m):.2f}, periods learned {statistics.mean(s["periods"] for *_, s in mine):.2f}')
    part3 = []
    for rank, n, robot, _ in rows:
        if rank == 0 and n > 1:
            r3 = torus.Robot(robot._room, 10_000_000)
            blind_robot.solution_trimmed(SimpleNamespace(move=r3.move), shortcuts=True)
            part3.append((robot.moves / (n - 1), r3.moves / (n - 1), robot.bumps == r3.bumps))
    print(f'  rank 0, Part 3\'s best variant: {statistics.mean(p for _, p, _ in part3):.2f} against '
          f'{statistics.mean(m for m, _, _ in part3):.2f}; same bumps on {sum(b for *_, b in part3)}'
          f'/{len(part3)}')


def report_no_limit(rng, rooms=100):
    """Frontier first with no radius at all, on rooms with a loop round the torus. Each room gets
    ten times the move calls the doubling robot needed there."""
    stuck = tried = 0
    while tried < rooms:
        room = torus.random_room(rng, rng.randint(3, 12), rng.randint(3, 12), rng.random() * 0.6)
        if not periods.periods(room):
            continue
        tried += 1
        safe = torus.Robot(room, 10_000_000)
        blind_torus.solution(safe)
        robot = torus.Robot(room, 10 * (safe.moves + safe.bumps))
        try:
            blind_torus.solution(robot, radius=math.inf)
        except RuntimeError:
            stuck += 1
    known = part3_map_size(CORRIDOR)
    print(f'\nFrontier first with no radius, {tried} rooms with a loop round the torus: still going '
          f"after ten times the doubling robot's move calls on {stuck}.")
    print(f"Part 3's robot on the corridor, after 20,000 move calls: its map holds {known:,} fields.")


def part3_map_size(room, limit=20_000):
    """How many fields Part 3's robot believes it has found when the limit stops it."""
    robot = torus.Robot(room, limit)
    pos, fields = (0, 0), {(0, 0)}
    step = {'^': (-1, 0), 'v': (1, 0), '<': (0, -1), '>': (0, 1)}
    try:
        blind_robot.solution(SimpleNamespace(move=robot.move))
    except RuntimeError:
        pass
    for d, ok in robot.log:
        if ok:
            pos = (pos[0] + step[d][0], pos[1] + step[d][1])
            fields.add(pos)
    return len(fields)


def report_growth(rng, rooms=400):
    """Why double? The same rooms with the radius growing 2, 3, 4 or 8 times."""
    sample = []
    while len(sample) < rooms:
        room = torus.random_room(rng, rng.randint(4, 14), rng.randint(4, 14), rng.random() * 0.55)
        if len(torus.floor(room)) >= 5:
            sample.append(room)
    print('\nGrowing the radius by other factors (moves per field, mean by rank, and the worst):')
    for growth in (2, 3, 4, 8):
        by = {0: [], 1: [], 2: []}
        for room in sample:
            robot = torus.Robot(room, 10_000_000)
            blind_torus.solution(robot, growth=growth)
            by[len(periods.periods(room))].append(robot.moves / (len(torus.floor(room)) - 1))
        means = ', '.join(f'rank {k} {statistics.mean(v):.2f}' for k, v in by.items() if v)
        print(f'  x{growth}: {means}; worst {max(max(v) for v in by.values() if v):.2f}')


if __name__ == '__main__':
    rng = random.Random(4)
    test_walk(rng)
    test_hermite(rng)
    test_prediction(rng)
    test_covers(rng)
    rows = test_landmark(rng)
    test_two_colour(rng)
    test_exact(rng)
    report(rows)
    report_no_limit(random.Random(9))
    report_growth(random.Random(42))
