"""Checks and a comparison report for emptying the bag (Part 3).

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

import dock_exact
import pick_cells_dock
import take_all_dock
from dirt_common import neighbours, parse
from dock_common import bfs_tree, bounds, split, split_reference, tree_order
from test_dirt import maze_room, random_room

STEP = {'^': (-1, 0), 'v': (1, 0), '<': (0, -1), '>': (0, 1)}
MIXES = ['12345', '5', '24', '35', '5555555551']
SOLVERS = {'pick cells': (pick_cells_dock.solution, True), 'take all': (take_all_dock.solution, False)}


def simulate_dock(room, cmds, capacity, pick):
    """Replay cmds and return the number of moves; fail on any broken rule."""
    dock, dirt = parse(room)
    assert set(cmds) <= set(STEP) | ({'C'} if pick else set()), f'unexpected command in {cmds!r}'
    pos, load, moves, cleaned = dock, 0, 0, set()
    for ch in cmds:
        if ch == 'C':
            assert pos != dock and dirt[pos] > 0 and pos not in cleaned, f'C on {pos}'
            assert load + dirt[pos] <= capacity, f'C at {pos} overflows the bag'
            load += dirt[pos]
            cleaned.add(pos)
            continue
        pos = (pos[0] + STEP[ch][0], pos[1] + STEP[ch][1])
        assert pos in dirt, f'walked into a wall at {pos}'
        moves += 1
        if pos == dock:
            load = 0
        elif not pick and dirt[pos] > 0 and pos not in cleaned:
            assert load + dirt[pos] <= capacity, f'entering {pos} overflows the bag'
            load += dirt[pos]
            cleaned.add(pos)
    assert cleaned == {f for f in dirt if dirt[f] > 0}, 'dirt left behind'
    lower, upper = bounds(dirt, bfs_tree(dock, dirt)[0], capacity)
    assert lower <= moves <= upper, f'{moves} moves outside [{lower}, {upper}]'
    return moves


def oracle_pick(room, capacity):
    """Pick cells by trips: cheapest trip for every set of fields (Held-Karp), then the best way
    to split all fields into closed trips plus one open trip."""
    dock, dirt = parse(room)
    fields = sorted(f for f in dirt if dirt[f] > 0)
    n = len(fields)
    if n == 0:
        return 0
    points = [dock] + fields
    dist = []
    for p in points:
        d, queue = {p: 0}, deque([p])
        while queue:
            cell = queue.popleft()
            for nb, _, _ in neighbours(cell, dirt):
                if nb not in d:
                    d[nb] = d[cell] + 1
                    queue.append(nb)
        dist.append([d[q] for q in points])
    inf = float('inf')
    size = 1 << n
    path = [[inf] * n for _ in range(size)]
    for j in range(n):
        path[1 << j][j] = dist[0][j + 1]
    for mask in range(1, size):
        for j in range(n):
            if path[mask][j] < inf:
                for k in range(n):
                    if not mask >> k & 1:
                        step = path[mask][j] + dist[j + 1][k + 1]
                        if step < path[mask | 1 << k][k]:
                            path[mask | 1 << k][k] = step
    load = [sum(dirt[fields[j]] for j in range(n) if mask >> j & 1) for mask in range(size)]
    closed = [min((path[m][j] + dist[j + 1][0] for j in range(n)), default=inf) if load[m] <= capacity else inf
              for m in range(size)]
    opened = [min(path[m], default=inf) if load[m] <= capacity else inf for m in range(size)]
    closed[0] = opened[0] = 0
    part = [0] + [inf] * (size - 1)
    for mask in range(1, size):
        low = mask & -mask
        sub = mask
        while sub:
            if sub & low:
                part[mask] = min(part[mask], closed[sub] + part[mask ^ sub])
            sub = (sub - 1) & mask
    full = size - 1
    best, sub = part[full], full
    while sub:
        best = min(best, opened[sub] + part[full ^ sub])
        sub = (sub - 1) & full
    return best


def greedy_cut(weights, d0, gaps, capacity):
    """Fill each trip until the next field doesn't fit (for comparison with the optimal cut)."""
    moves, start, load = 0, 0, 0
    for k, w in enumerate(weights):
        if load + w > capacity:
            moves += d0[start] + sum(gaps[start:k - 1]) + d0[k - 1]
            start, load = k, 0
        load += w
    return moves + d0[start] + sum(gaps[start:len(weights) - 1])


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


def solve(room, capacity, rule):
    solver, pick = SOLVERS[rule]
    return simulate_dock(room, solver(room, capacity), capacity, pick)


def rooms_from_post():
    corridor = ['#######', '#*3322#', '#######']
    blocked = ['######', '#*154#', '######']
    fives = ['#####', '#*55#', '#555#', '#####']
    for name, room, pick_best, take_best in (('3322', corridor, 10, 10), ('154', blocked, 7, 9),
                                             ('fives', fives, 15, 15)):
        for rule, wanted in (('pick cells', pick_best), ('take all', take_best)):
            exact = simulate_dock(room, dock_exact.solution(room, 5, rule == 'pick cells'), 5, rule == 'pick cells')
            expect(f'{name} exact {rule}', exact, wanted)
            expect(f'{name} {rule}', solve(room, 5, rule), wanted)
        dock, dirt = parse(room)
        dist, parent = bfs_tree(dock, dirt)
        lower, upper = bounds(dirt, dist, 5)
        print(f'  {name:<6} pick cells {pick_best}, take all {take_best}; lower bound {lower}, one field per trip {upper}')
    dock, dirt = parse(corridor)
    dist, parent = bfs_tree(dock, dirt)
    order = tree_order(dock, dirt, dist, parent)
    gaps = [1] * (len(order) - 1)
    weights, d0 = [dirt[f] for f in order], [dist[f] for f in order]
    expect('3322 greedy cut', greedy_cut(weights, d0, gaps, 5), 12)
    expect('3322 optimal cut', split(weights, d0, gaps, 5)[0], 10)
    print('rooms from the post ..... ok (3322: filling each trip gives 12, the optimal cut 10)')


def split_checks(rng, count=2000):
    for _ in range(count):
        n = rng.randint(1, 12)
        weights = [rng.randint(1, 5) for _ in range(n)]
        d0 = [rng.randint(1, 9) for _ in range(n)]
        gaps = [rng.randint(1, 6) for _ in range(n - 1)]
        capacity = rng.randint(5, 15)
        moves, trips = split(weights, d0, gaps, capacity)
        expect('split vs reference', moves, split_reference(weights, d0, gaps, capacity))
        assert [i for trip in trips for i in trip] == list(range(n))
        assert all(sum(weights[i] for i in trip) <= capacity for trip in trips)
        cost = sum(d0[t[0]] + sum(gaps[t[0]:t[-1]]) + d0[t[-1]] for t in trips) - d0[trips[-1][-1]]
        expect('split cost of its own trips', cost, moves)
    print(f'optimal cut ............. matches the plain double loop on {count} random inputs')


def tiny_rooms(rng, count=400, max_fields=9):
    stats = {rule: [0, 0] for rule in SOLVERS}          # [best, worst gap]
    checked = 0
    while checked < count:
        R, C, mix = rng.randint(3, 6), rng.randint(3, 6), rng.choice(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.4]), mix)
        _, dirt = parse(room)
        if sum(1 for f in dirt if dirt[f] > 0) > max_fields:
            continue
        capacity = rng.choice([5, 6, 7, 8, 10, 15, max(5, sum(dirt.values()) + 1)])
        exact = {}
        for rule, (_, pick) in SOLVERS.items():
            exact[rule] = simulate_dock(room, dock_exact.solution(room, capacity, pick), capacity, pick)
            got = solve(room, capacity, rule)
            assert got >= exact[rule], f'{rule} beat the exact search on {room}'
            stats[rule][0] += got == exact[rule]
            stats[rule][1] = max(stats[rule][1], got - exact[rule])
        expect(f'trip oracle {room} K={capacity}', oracle_pick(room, capacity), exact['pick cells'])
        assert exact['pick cells'] <= exact['take all']
        checked += 1
    print(f'tiny rooms (<= {max_fields} dirty) . {checked} rooms: the exact search and the trip oracle agree on every room')
    for rule, (best, gap) in stats.items():
        print(f'  {rule:<22}  best on {best}/{checked}, worst gap {gap}')


def exact_cases(rng):
    for _ in range(20):
        room = random_room(rng, rng.randint(4, 12), rng.randint(4, 12), 0.2, '5')
        dock, dirt = parse(room)
        upper = bounds(dirt, bfs_tree(dock, dirt)[0], 5)[1]
        for rule in SOLVERS:
            expect(f'all fives, {rule}', solve(room, rng.randint(5, 9), rule), upper)
    for _ in range(20):
        room = maze_room(rng, 11, 11, '12345')
        dock, dirt = parse(room)
        dist = bfs_tree(dock, dirt)[0]
        wanted = 2 * (len(dirt) - 1) - max(dist.values())
        for rule in SOLVERS:
            expect(f'maze with a big bag, {rule}', solve(room, sum(dirt.values()), rule), wanted)
    print('exact cases ............. all 5s with capacity 5-9 = one field per trip; mazes with a big bag = 2(N-1) - farthest')


def big_rooms(rng, count=8):
    rows = {rule: [] for rule in SOLVERS}
    for i in range(count):
        mix = MIXES[i % len(MIXES)]
        room = random_room(rng, 40, 40, rng.choice([0, 0.15, 0.3]), mix)
        capacity = rng.choice([5, 10, 20, 50])
        dock, dirt = parse(room)
        lower, _ = bounds(dirt, bfs_tree(dock, dirt)[0], capacity)
        for rule in SOLVERS:
            t = time.perf_counter()
            moves = solve(room, capacity, rule)
            rows[rule].append((moves / max(1, lower), time.perf_counter() - t))
    empty = ['#' * 40] + ['#*' + '5' * 37 + '#'] + ['#' + '5' * 38 + '#'] * 37 + ['#' * 40]
    worst = solve(empty, 5, 'take all')
    print(f'big rooms (40x40) ....... {count} rooms per rule; worst case (all 5s, capacity 5): {worst:,} moves')
    for rule, data in rows.items():
        print(f'  {rule:<22}  moves / lower bound avg {sum(r for r, _ in data) / count:.2f}, '
              f'time avg {sum(t for _, t in data) / count:.1f} s, worst {max(t for _, t in data):.1f} s')


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


if __name__ == '__main__':
    main()
