"""Checks and a comparison report for the fewest-commands solutions (Part 3).

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

import shortest_route_exact
import shortest_route_fast
from cleaning_robot import solution_trimmed as part1_trimmed
from shortest_route_fast import lower_bound, parse, steps
from test_cleaning_robot import EXAMPLES, empty_room, random_room, simulate
from test_dirt import maze_room

CORRIDOR = ['#########', '#.....*.#', '#########']
TWO_LOOPS = ['#####', '#..##', '#.*.#', '##..#', '#####']


def expect(name, got, wanted):
    assert got == wanted, f'{name}: got {got}, expected {wanted}'


def distances(start, cells):
    dist, queue = {start: 0}, deque([start])
    while queue:
        cell = queue.popleft()
        for n, _ in steps(cell, cells):
            if n not in dist:
                dist[n] = dist[cell] + 1
                queue.append(n)
    return dist


def check_route(room, route):
    """The route must follow the rules and clean every floor field."""
    _, cells = parse(room)
    expect('cleaned fields', simulate(room, route), cells)
    return len(route)


# ---------- brute-force referees (tiny rooms) ----------

def referee_bfs(room):
    """Fewest commands by plain breadth-first search over (field, cleaned fields)."""
    start, cells = parse(room)
    if len(cells) == 1:
        return 0
    first = (start, frozenset([start]))
    seen, layer, k = {first}, [first], 0
    while True:
        k += 1
        nxt = []
        for pos, done in layer:
            for n, _ in steps(pos, cells):
                state = (n, done | {n})
                if len(state[1]) == len(cells):
                    return k
                if state not in seen:
                    seen.add(state)
                    nxt.append(state)
        layer = nxt


def referee_orders(room):
    """Fewest commands as the cheapest visiting order joined by shortest paths (Held-Karp)."""
    start, cells = parse(room)
    order = [start] + sorted(cells - {start})
    n = len(order)
    rows = [distances(a, cells) for a in order]
    dist = [[row[b] for b in order] for row in rows]
    inf = float('inf')
    best = [[inf] * n for _ in range(1 << n)]
    best[1][0] = 0
    for mask in range(1, 1 << n, 2):
        for j in range(n):
            d = best[mask][j]
            if d == inf:
                continue
            for k in range(n):
                if not mask >> k & 1 and d + dist[j][k] < best[mask | 1 << k][k]:
                    best[mask | 1 << k][k] = d + dist[j][k]
    return min(best[(1 << n) - 1])


def longest_path(start, cells):
    """Length of the longest simple path from start (the spanning-tree formula's L)."""
    best, stack = 0, [(start, frozenset([start]))]
    while stack:
        cell, seen = stack.pop()
        best = max(best, len(seen) - 1)
        for n, _ in steps(cell, cells):
            if n not in seen:
                stack.append((n, seen | {n}))
    return best


# ---------- checks ----------

def rooms_from_post():
    for name, room, best, trimmed in (('corridor', CORRIDOR, 7, 11), ('two loops', TWO_LOOPS, 7, 9)):
        fast = check_route(room, shortest_route_fast.solution(room))
        exact = check_route(room, shortest_route_exact.solution(room))
        expect(f'{name} exact', exact, best)
        expect(f'{name} lower bound', lower_bound(room), best)
        expect(f'{name} Part 1 trimmed', len(part1_trimmed(room)), trimmed)
        print(f'  {name:<10} Part 1 trimmed {trimmed}, fast {fast}, exact {exact}, lower bound {best}')
    room = EXAMPLES[1][0]
    fast = check_route(room, shortest_route_fast.solution(room))
    expect('example 2 fast reaches the lower bound', fast, lower_bound(room))
    print(f'  example 2  Part 1 trimmed {len(part1_trimmed(room))}, fast {fast} = lower bound {lower_bound(room)}')
    print('rooms from the post ..... ok')


def tiny_rooms(rng, count=800, max_cells=12):
    checked = matches = tight = beats_formula = 0
    worst_gap, smallest_beat = 0, None
    while checked < count:
        R, C = rng.randint(3, 6), rng.randint(3, 6)
        room = (maze_room(rng, R, C, '.') if rng.random() < 0.2
                else random_room(rng, R, C, rng.choice([0, 0.2, 0.4])))
        start, cells = parse(room)
        n = len(cells)
        if n > max_cells:
            continue
        best = referee_bfs(room)
        if n <= 9:
            expect(f'visiting-order referee {room}', referee_orders(room), best)
        exact = shortest_route_exact.solution(room)
        assert exact is not None, room
        expect(f'exact {room}', check_route(room, exact), best)

        fast = check_route(room, shortest_route_fast.solution(room))
        ecc = max(distances(start, cells).values())
        bound = lower_bound(room)
        assert bound <= best <= fast <= min(len(part1_trimmed(room)), 2 * (n - 1) - ecc), room
        matches += fast == best
        worst_gap = max(worst_gap, fast - best)
        tight += bound == best
        if best < 2 * (n - 1) - longest_path(start, cells):
            beats_formula += 1
            if smallest_beat is None or n < len(parse(smallest_beat)[1]):
                smallest_beat = room
        checked += 1
    print(f'tiny rooms (<= {max_cells} fields) {checked} rooms: exact = both brute-force referees on every room')
    print(f'  fast .................. best on {matches}/{checked}, worst gap {worst_gap}')
    print(f'  lower bound ........... tight on {tight}/{checked}')
    print(f'  tree formula .......... beaten on {beats_formula} rooms; smallest:')
    for row in smallest_beat or []:
        print(f'      {row}')


def small_rooms(rng, count=200, max_states=200_000):
    same = shorter = gave_up = 0
    fast_time = exact_time = 0.0
    for _ in range(count):
        R, C = rng.randint(6, 8), rng.randint(6, 8)
        room = (maze_room(rng, R, C, '.') if rng.random() < 0.2
                else random_room(rng, R, C, rng.choice([0, 0.2, 0.35])))
        t = time.perf_counter()
        fast = check_route(room, shortest_route_fast.solution(room))
        fast_time += time.perf_counter() - t
        t = time.perf_counter()
        exact = shortest_route_exact.solution(room, max_states)
        exact_time += time.perf_counter() - t
        if exact is None:
            gave_up += 1
            continue
        exact = check_route(room, exact)
        assert exact <= fast
        same += exact == fast
        shorter += exact < fast
    print(f'small rooms (6x6 to 8x8)  {count} rooms: same {same}, exact shorter {shorter}, '
          f'exact gave up {gave_up}')
    print(f'  time .................. fast avg {1000 * fast_time / count:.1f} ms, '
          f'exact avg {1000 * exact_time / count:.0f} ms')


def big_rooms(rng, count=20):
    proven = mazes = 0
    gaps, saved, times = [], [], []
    for i in range(count):
        maze = i % 4 == 0
        room = maze_room(rng, 39, 39, '.') if maze else random_room(rng, 40, 40, rng.choice([0, 0.15, 0.3]))
        start, cells = parse(room)
        n = len(cells)
        t = time.perf_counter()
        fast = check_route(room, shortest_route_fast.solution(room))
        times.append(time.perf_counter() - t)
        if maze:
            expect('maze route', fast, 2 * (n - 1) - max(distances(start, cells).values()))
            mazes += 1
        gaps.append(fast - lower_bound(room))
        proven += gaps[-1] == 0
        saved.append(1 - fast / len(part1_trimmed(room)))
    print(f'big rooms (40x40) ....... {count} rooms ({mazes} mazes, each exactly 2(N-1) - ecc): '
          f'proven best {proven}, worst gap to bound {max(gaps)}')
    print(f'  vs Part 1 trimmed ..... {100 * sum(saved) / count:.0f}% shorter on average; '
          f'time avg {1000 * sum(times) / count:.0f} ms, worst {1000 * max(times):.0f} ms')


def empty_rectangles(max_size=40):
    misses = []
    for R in range(3, max_size + 1):
        for C in range(3, max_size + 1):
            room = empty_room(R, C, (1, 1))
            if len(shortest_route_fast.solution(room)) != (R - 2) * (C - 2) - 1:
                misses.append((R, C))
    print(f'empty rooms from a corner  every size 3x3 to {max_size}x{max_size}: '
          f'N-1 commands on all but {len(misses)} {misses[:5]}')


def main():
    rng = random.Random(2026)
    rooms_from_post()
    tiny_rooms(rng)
    small_rooms(rng)
    big_rooms(rng)
    empty_rectangles()
    print('all checks passed')


if __name__ == '__main__':
    main()
