"""Take all: every cell the robot enters is emptied, so it may only enter cells that still fit.

Fast and near-best: never more than 4 below the best possible, and usually proven best.
"""
import heapq
import random

from dirt_common import best_subset, neighbours, parse, take_all_route


def spanning_tree(start, dirt, keep, rng):
    """Random spanning tree grown from start (Prim). Cells in `keep` are attached before any
    other cell, so a connected `keep` containing start becomes a subtree hanging from the root."""
    parent, children = {start: None}, {start: []}
    heap = []
    cell = start
    while True:
        for n, _, _ in neighbours(cell, dirt):
            if n not in parent:
                heapq.heappush(heap, (n not in keep, rng.random(), n, cell))
        while heap and heap[0][2] in parent:
            heapq.heappop(heap)
        if not heap:
            return parent, children
        _, _, cell, par = heapq.heappop(heap)
        parent[cell] = par
        children[cell] = []
        children[par].append(cell)


def best_subtree(start, parent, children, dirt, capacity):
    """Exact best set that is connected inside this tree, contains start and fits capacity."""
    order, stack = [], [start]
    while stack:                                # preorder: each subtree is one contiguous block
        cell = stack.pop()
        order.append(cell)
        stack.extend(children[cell])
    size = {cell: 1 for cell in order}
    for cell in reversed(order[1:]):
        size[parent[cell]] += size[cell]

    n = len(order)
    w = [dirt[cell] for cell in order]
    span = [size[cell] for cell in order]
    mask = (1 << (capacity + 1)) - 1
    # dp[i]: bit s is set when the choices for order[0..i-1] collect exactly s
    # and order[i] may still be taken (its parent is in the set).
    dp = [0] * (n + 1)
    dp[1] = 1                                   # the start is always in the set, with no dirt
    for i in range(1, n):
        if dp[i]:
            dp[i + 1] |= (dp[i] << w[i]) & mask     # take order[i]
            dp[i + span[i]] |= dp[i]                # skip order[i] and its whole subtree

    best = dp[n].bit_length() - 1
    skips_into = [[] for _ in range(n + 1)]
    for i in range(1, n):
        skips_into[i + span[i]].append(i)
    chosen, i, s = {start}, n, best
    while i > 1:
        if s >= w[i - 1] and (dp[i - 1] >> (s - w[i - 1])) & 1:
            chosen.add(order[i - 1])
            s -= w[i - 1]
            i -= 1
        else:
            i = next(k for k in skips_into[i] if (dp[k] >> s) & 1)
    return best, chosen


def best_connected_set(start, dirt, capacity, rounds=30, seed=0):
    """Return (best, cells, upper).

    cells is a connected set containing start whose dirt, best, fits capacity.
    upper is a bound no connected set can beat, so best == upper proves best is optimal.
    """
    total = sum(dirt.values())
    if total <= capacity:
        return total, set(dirt), total
    upper, _ = best_subset(list(dirt.values()), capacity)
    rng = random.Random(seed)
    best, cells = 0, {start}
    for round_no in range(rounds):
        if best == upper:
            break
        # Even rounds build the tree around the best set, so they can only improve on it.
        # Odd rounds use a fully random tree, which can reshape the set out of a dead end.
        keep = cells if round_no % 2 == 0 else {start}
        parent, children = spanning_tree(start, dirt, keep, rng)
        value, chosen = best_subtree(start, parent, children, dirt, capacity)
        if value >= best:                       # on a tie, move to the new set to keep exploring
            best, cells = value, chosen
    return best, cells, upper


def solution(room, capacity):
    start, dirt = parse(room)
    _, cells, _ = best_connected_set(start, dirt, capacity)
    return take_all_route(start, cells)
