"""Checks and a comparison report for the dirty-room robot solutions.

Run from this folder: python test_dirt.py
"""
import itertools
import random
import time
from functools import lru_cache

import pick_cells
import pick_cells_battery
import take_all_battery
import take_all_exact
import take_all_fast
from dirt_common import neighbours, parse
from take_all_fast import best_connected_set

STEP = {'^': (-1, 0), 'v': (1, 0), '<': (0, -1), '>': (0, 1)}
DIRT_MIXES = ['12345', '5', '24', '35', '5555555551']


# ---------- simulator ----------

def simulate(room, cmds, capacity, battery=None, pick=False):
    """Replay cmds and return the dirt collected; fail on any broken rule."""
    start, dirt = parse(room)
    assert len(cmds) <= 50_000, 'more than 50,000 commands'
    allowed = set(STEP) | ({'C'} if pick else set())
    assert set(cmds) <= allowed, f'unexpected command in {cmds!r}'
    moves = sum(ch in STEP for ch in cmds)
    assert battery is None or moves <= battery, f'{moves} moves, battery {battery}'
    pos, collected, cleaned = start, 0, {start}
    for ch in cmds:
        if ch == 'C':
            assert pos not in cleaned and dirt[pos] > 0, f'C on a clean cell {pos}'
            assert collected + dirt[pos] <= capacity, f'C at {pos} overflows capacity'
            collected += dirt[pos]
            cleaned.add(pos)
            continue
        dr, dc = STEP[ch]
        pos = (pos[0] + dr, pos[1] + dc)
        assert pos in dirt, f'walked into a wall at {pos}'
        if not pick and pos not in cleaned:
            assert collected + dirt[pos] <= capacity, f'entering {pos} overflows capacity'
            collected += dirt[pos]
            cleaned.add(pos)
    return collected


# ---------- room generators ----------

def finish_room(rng, grid):
    """Put the start on a random floor cell and wall off floor it can't reach."""
    floor = [(r, c) for r, row in enumerate(grid) for c, ch in enumerate(row) if ch != '#']
    if not floor:
        grid[1][1] = '1'
        floor = [(1, 1)]
    start = rng.choice(floor)
    reach, stack = {start}, [start]
    while stack:
        r, c = stack.pop()
        for n in ((r - 1, c), (r + 1, c), (r, c - 1), (r, c + 1)):
            if grid[n[0]][n[1]] != '#' and n not in reach:
                reach.add(n)
                stack.append(n)
    for r, c in floor:
        if (r, c) not in reach:
            grid[r][c] = '#'
    grid[start[0]][start[1]] = '*'
    return [''.join(row) for row in grid]


def random_room(rng, R, C, wall_p, mix):
    grid = [['#'] * C for _ in range(R)]
    for r in range(1, R - 1):
        for c in range(1, C - 1):
            if rng.random() >= wall_p:
                grid[r][c] = rng.choice(mix)
    return finish_room(rng, grid)


def maze_room(rng, R, C, mix):
    """A perfect maze: its floor has no loops, so the fast solution must be exact on it."""
    grid = [['#'] * C for _ in range(R)]
    grid[1][1] = rng.choice(mix)
    stack = [(1, 1)]
    while stack:
        r, c = stack[-1]
        options = [(r + dr, c + dc) for dr, dc in ((-2, 0), (2, 0), (0, -2), (0, 2))
                   if 0 < r + dr < R - 1 and 0 < c + dc < C - 1 and grid[r + dr][c + dc] == '#']
        if not options:
            stack.pop()
            continue
        nr, nc = rng.choice(options)
        grid[(r + nr) // 2][(c + nc) // 2] = rng.choice(mix)
        grid[nr][nc] = rng.choice(mix)
        stack.append((nr, nc))
    return finish_room(rng, grid)


def pick_capacity(rng, total):
    return rng.choice([1, rng.randint(1, max(1, total)), max(1, total - 1), max(1, total), total + 5])


# ---------- brute-force oracles (tiny rooms only) ----------

def oracle_take_all(room, capacity):
    """Best dirt over every connected set containing the start."""
    start, dirt = parse(room)
    first = frozenset([start])
    best, seen, stack = 0, {first}, [(first, 0)]
    while stack:
        cells, total = stack.pop()
        best = max(best, total)
        for cell in cells:
            for n, _, _ in neighbours(cell, dirt):
                if n not in cells and total + dirt[n] <= capacity:
                    bigger = cells | {n}
                    if bigger not in seen:
                        seen.add(bigger)
                        stack.append((bigger, total + dirt[n]))
    return best


def oracle_pick(room, capacity):
    """Best dirt over every choice of how many 1s, 2s, ... 5s to clean."""
    _, dirt = parse(room)
    counts = [sum(v == w for v in dirt.values()) for w in range(1, 6)]
    best = 0
    for picks in itertools.product(*(range(c + 1) for c in counts)):
        s = sum((w + 1) * k for w, k in enumerate(picks))
        if s <= capacity:
            best = max(best, s)
    return best


def oracle_battery(room, capacity, battery, pick):
    """Best dirt over every walk of at most `battery` moves."""
    start, dirt = parse(room)
    cells = list(dirt)
    bit = {cell: 1 << i for i, cell in enumerate(cells)}

    @lru_cache(maxsize=None)
    def value(mask):
        weights = [dirt[cell] for cell in cells if mask & bit[cell]]
        if not pick:
            return sum(weights)
        sums = {0}
        for w in weights:
            sums |= {s + w for s in sums if s + w <= capacity}
        return max(sums)

    layer = {(start, bit[start])}
    seen, best = set(layer), 0
    for _ in range(battery):
        nxt = set()
        for pos, mask in layer:
            for n, _, _ in neighbours(pos, dirt):
                m = mask | bit[n]
                if not pick and m != mask and value(m) > capacity:
                    continue
                if (n, m) not in seen:
                    seen.add((n, m))
                    nxt.add((n, m))
                    best = max(best, value(m))
        if not nxt:
            break
        layer = nxt
    return best


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

def timed(fn, *args):
    t = time.perf_counter()
    out = fn(*args)
    return out, time.perf_counter() - t


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


def regression_rooms():
    corridor = ['######', '#3*51#', '######']      # greedy would take the 3 and stop
    fives = ['#####', '#*55#', '#555#', '#####']
    too_dirty = ['####', '#*3#', '#44#', '####']
    lonely = ['###', '#*#', '###']
    for solve in (take_all_fast.solution, take_all_exact.solution):
        expect('corridor', simulate(corridor, solve(corridor, 6), 6), 6)
        expect('all fives', simulate(fives, solve(fives, 12), 12), 10)
        expect('too dirty', solve(too_dirty, 2), '')
        expect('start only', solve(lonely, 5), '')
    expect('pick corridor', simulate(corridor, pick_cells.solution(corridor, 6), 6, pick=True), 6)
    expect('pick fives', simulate(fives, pick_cells.solution(fives, 12), 12, pick=True), 10)
    expect('pick too dirty', pick_cells.solution(too_dirty, 2), '')
    expect('pick start only', pick_cells.solution(lonely, 5), '')
    expect('take-all battery 0', take_all_battery.solution(fives, 12, 0), '')
    expect('pick battery 0', pick_cells_battery.solution(fives, 12, 0), '')
    print('regression rooms ........ ok')


def tiny_rooms(rng, count=800, max_cells=12):
    checked = 0
    fast_misses = []
    ratios = {'take_all_battery': [], 'pick_cells_battery': []}
    while checked < count:
        R, C, mix = rng.randint(3, 6), rng.randint(3, 6), rng.choice(DIRT_MIXES)
        room = maze_room(rng, R, C, mix) if rng.random() < 0.2 else random_room(rng, R, C, rng.choice([0, 0.2, 0.4]), mix)
        _, dirt = parse(room)
        n, total = len(dirt), sum(dirt.values())
        if n > max_cells:
            continue
        K = pick_capacity(rng, total)
        best_take, best_pick = oracle_take_all(room, K), oracle_pick(room, K)
        if checked % 10 == 0:
            expect(f'oracles agree (take all) {room} K={K}', oracle_battery(room, K, 4 * n, False), best_take)
            expect(f'oracles agree (pick) {room} K={K}', oracle_battery(room, K, 4 * n, True), best_pick)

        expect(f'pick_cells {room} K={K}', simulate(room, pick_cells.solution(room, K), K, pick=True), best_pick)
        expect(f'take_all_exact {room} K={K}', simulate(room, take_all_exact.solution(room, K), K), best_take)
        fast = simulate(room, take_all_fast.solution(room, K), K)
        assert min(total, K - 4) <= fast <= best_take, f'take_all_fast {room} K={K}: {fast} vs {best_take}'
        if fast < best_take:
            fast_misses.append((room, K, fast, best_take))

        B = rng.randint(0, 2 * n)
        for name, module, pick in (('take_all_battery', take_all_battery, False),
                                   ('pick_cells_battery', pick_cells_battery, True)):
            got = simulate(room, module.solution(room, K, B), K, battery=B, pick=pick)
            best = oracle_battery(room, K, B, pick)
            assert got <= best, f'{name} beat the oracle?! {room} K={K} B={B}'
            ratios[name].append(got / best if best else 1.0)
        checked += 1

    print(f'tiny rooms (<= {max_cells} cells) . {checked} rooms, all outputs valid')
    print('  pick_cells ............ equals brute force on every room')
    print('  take_all_exact ........ equals brute force on every room')
    print(f'  take_all_fast ......... best on {checked - len(fast_misses)}/{checked} rooms')
    for room, K, fast, best in fast_misses[:1]:
        print(f'    a room it missed: capacity {K}, fast {fast}, best {best}')
        for row in room:
            print(f'      {row}')
    for name, rs in ratios.items():
        best_count = sum(r == 1.0 for r in rs)
        print(f'  {name:<22}  best on {best_count}/{len(rs)}, average {sum(rs) / len(rs):.1%} of best, worst {min(rs):.1%}')


def compare_fast_exact(rng, count=300):
    proven = matches = 0
    gaps, fast_times, exact_times = [], [], []
    for _ in range(count):
        R, C, mix = rng.randint(5, 8), rng.randint(5, 8), rng.choice(DIRT_MIXES)
        room = maze_room(rng, R, C, mix) if rng.random() < 0.25 else random_room(rng, R, C, rng.choice([0, 0.2, 0.35]), mix)
        start, dirt = parse(room)
        K = pick_capacity(rng, sum(dirt.values()))
        (fast, _, upper), t_fast = timed(best_connected_set, start, dirt, K)
        (exact, _), t_exact = timed(take_all_exact.best_connected_set_exact, start, dirt, K)
        assert fast <= exact <= upper
        proven += fast == upper
        matches += fast == exact
        gaps.append(exact - fast)
        fast_times.append(t_fast)
        exact_times.append(t_exact)
    print(f'fast vs exact (take all) . {count} rooms up to 8x8')
    print(f'  same answer ........... {matches}/{count}  (fast proved itself best on {proven})')
    print(f'  gap when different .... average {sum(gaps) / max(1, count - matches):.2f}, worst {max(gaps)}')
    print(f'  time .................. fast avg {1000 * sum(fast_times) / count:.1f} ms, worst {1000 * max(fast_times):.0f} ms;'
          f' exact avg {1000 * sum(exact_times) / count:.1f} ms, worst {1000 * max(exact_times):.0f} ms')


def big_rooms(rng, count=20):
    proven = 0
    times = {name: [] for name in ('take_all_fast', 'pick_cells', 'take_all_battery', 'pick_cells_battery')}
    for i in range(count):
        mix = DIRT_MIXES[i % len(DIRT_MIXES)]
        room = maze_room(rng, 39, 39, mix) if i % 4 == 0 else random_room(rng, 40, 40, rng.choice([0, 0.15, 0.3]), mix)
        start, dirt = parse(room)
        total = sum(dirt.values())
        K, B = pick_capacity(rng, total), rng.randint(0, 3000)
        best, _, upper = best_connected_set(start, dirt, K)
        proven += best == upper

        route, t = timed(take_all_fast.solution, room, K)
        times['take_all_fast'].append(t)
        got = simulate(room, route, K)
        assert got == best and min(total, K - 4) <= got <= upper

        route, t = timed(pick_cells.solution, room, K)
        times['pick_cells'].append(t)
        expect('pick_cells big room', simulate(room, route, K, pick=True), upper)

        route, t = timed(take_all_battery.solution, room, K, B)
        times['take_all_battery'].append(t)
        assert simulate(room, route, K, battery=B) <= upper

        route, t = timed(pick_cells_battery.solution, room, K, B)
        times['pick_cells_battery'].append(t)
        assert simulate(room, route, K, battery=B, pick=True) <= upper

    print(f'big rooms (40x40) ....... {count} rooms, all outputs valid; take_all_fast proved itself best on {proven}')
    for name, ts in times.items():
        print(f'  {name:<22}  avg {1000 * sum(ts) / len(ts):.0f} ms, worst {1000 * max(ts):.0f} ms')


def main():
    rng = random.Random(2026)
    regression_rooms()
    tiny_rooms(rng)
    compare_fast_exact(rng)
    big_rooms(rng)
    print('all checks passed')


if __name__ == '__main__':
    main()
