"""Part 9: the fewest commands, exactly, on big rooms: the integer program in HiGHS.

branch_cut.py solves the program from scratch and is only fit for small rooms. Here HiGHS solves
it, with every trick a modern solver has: presolve, its own cuts, heuristics and a parallel
branch-and-bound tree. Two things are ours:

- Part 8's cutting-plane loop runs first, with HiGHS as the linear solver, and every connected
  rule it found goes into the integer program. They are what make its bound strong.
- The connected rules can't all be written down, and highspy can't yet add rules from inside the
  search, so connectivity is also said a second, compact way: the start sends one unit of a
  commodity to every other field, along pairs of neighbours the route walks, at most N - 1 over
  each. A set of walks that doesn't connect every field can't carry that flow. It adds only
  2 variables per pair and one rule per field and per pair, so it can all be written down; alone
  its bound is weak, but the connected rules make up for that.

Part 3's fast route goes in as the first answer, so HiGHS starts with something to beat. With a
time limit the result may not be proven best: then `bound` says how far off it can be.
"""
import math
import time
from collections import deque

import route_lp
from branch_cut import euler_route
from shortest_route_fast import solution as fast_route


def route_counts(model, route):
    """A route's walks per pair of neighbours (any walked more than twice brought down to 1 or 2,
    which keeps parity and connectivity) and where it ends."""
    index = {f: i for i, f in enumerate(model.fields)}
    edge = {frozenset(e): k for k, e in enumerate(model.edges)}
    step = {'^': (-1, 0), 'v': (1, 0), '<': (0, -1), '>': (0, 1)}
    x = [0] * model.m
    r, c = model.fields[model.s]
    for ch in route:
        dr, dc = step[ch]
        x[edge[frozenset((index[(r, c)], index[(r + dr, c + dc)]))]] += 1
        r, c = r + dr, c + dc
    x = [v if v <= 2 else 2 - v % 2 for v in x]
    z = [0] * model.n
    z[index[(r, c)]] = 1
    return x, z


def tree_flow(model, x):
    """A flow for walks x: along a breadth-first tree of the walked pairs, each arc carries one unit
    for every field beyond it. Arc 2k runs along edge k as listed, 2k + 1 against it."""
    n, s = model.n, model.s
    parent = {s: None}
    order, queue = [s], deque([s])
    while queue:
        u = queue.popleft()
        for k in model.touching[u]:
            if x[k]:
                a, b = model.edges[k]
                v = b if a == u else a
                if v not in parent:
                    parent[v] = (u, k)
                    order.append(v)
                    queue.append(v)
    assert len(order) == n, 'the walks do not reach every field'
    below = [1] * n
    f = [0.0] * (2 * model.m)
    for v in reversed(order[1:]):
        u, k = parent[v]
        below[u] += below[v]
        f[2 * k if model.edges[k][0] == u else 2 * k + 1] = float(below[v])
    return f


def solve(room, time_limit=600, threads=None, cuts=True, start=True, relax=False):
    """Returns a dict: route (the best found), bound (no route is shorter), optimal (proven),
    lp (Part 8's bound), nodes, seconds.

    cuts=False leaves out Part 8's connected rules, start=False leaves out the fast route; they are
    there to measure what each is worth. relax=True drops the whole-number rules and returns only
    the linear program's value, as 'relaxation'."""
    import highspy
    import numpy as np
    t0 = time.perf_counter()
    lp = route_lp.lp_bound(room, 'highs')
    model = lp.model
    n, m, s = model.n, model.m, model.s
    if n == 1:
        return {'route': '', 'bound': 0, 'optimal': True, 'lp': 0, 'nodes': 0, 'seconds': 0.0}
    h = highspy.Highs()
    h.setOptionValue('output_flag', False)
    h.setOptionValue('time_limit', float(time_limit))
    h.setOptionValue('mip_rel_gap', 0.0)
    if threads:
        h.setOptionValue('threads', threads)
    # columns: x (m, whole 0..2), z (n, whole 0..1), w (n, whole 0..4), flow (2m, 0..n-1)
    X, Z, W, F = 0, m, m + n, m + 2 * n
    count = m + 2 * n + 2 * m
    h.addVars(count, np.zeros(count), np.r_[np.full(m, 2.0), np.ones(n), np.full(n, 4.0), np.full(2 * m, n - 1.0)])
    h.changeColsCost(m, np.arange(m, dtype=np.int32), np.ones(m))
    whole = m + 2 * n
    if not relax:
        h.changeColsIntegrality(whole, np.arange(whole, dtype=np.int32), np.array([highspy.HighsVarType.kInteger] * whole))

    def row(low, high, coef):
        h.addRow(float(low), float(high), len(coef), np.array(list(coef), dtype=np.int32),
                 np.array(list(coef.values()), dtype=float))
    row(1, 1, {Z + v: 1 for v in range(n)})
    for v in range(n):                                  # parity
        coef = {X + k: 1 for k in model.touching[v]}
        coef[Z + v] = -1 if v == s else 1
        coef[W + v] = -2
        row(-1 if v == s else 0, -1 if v == s else 0, coef)
    for v in range(n):                                  # the flow: n - 1 units out of s, 1 into every other field
        coef = {}
        for k in model.touching[v]:
            out, back = (F + 2 * k, F + 2 * k + 1) if model.edges[k][0] == v else (F + 2 * k + 1, F + 2 * k)
            coef[out] = coef.get(out, 0) + 1
            coef[back] = coef.get(back, 0) - 1
        row(n - 1 if v == s else -1, n - 1 if v == s else -1, coef)
    for k in range(m):                                  # only over walked pairs
        row(-np.inf, 0, {F + 2 * k: 1, F + 2 * k + 1: 1, X + k: -(n - 1)})
    for S in (lp.cuts if cuts else []):                 # Part 8's connected rules
        row(2, np.inf, model.cut_row(S))
    if relax:
        h.run()
        return {'relaxation': h.getInfo().objective_function_value, 'lp': lp.value}
    # Part 3's fast route as the first answer
    x, z = route_counts(model, fast_route(room))
    w = [(sum(x[k] for k in model.touching[v]) + (-z[v] + 1 if v == s else z[v])) // 2 for v in range(n)]
    first = highspy.HighsSolution()
    first.col_value = [float(v) for v in x + z + w] + tree_flow(model, x)
    first.value_valid = True
    if start:
        h.setSolution(first)
    h.run()
    info = h.getInfo()
    values = list(h.getSolution().col_value)
    x = [round(v) for v in values[X:X + m]]
    z = [round(v) for v in values[Z:Z + n]]
    route = euler_route(model, x, z)
    optimal = h.getModelStatus() == highspy.HighsModelStatus.kOptimal
    bound = max(lp.commands, math.ceil(info.mip_dual_bound - 1e-6))  # a route has a whole number of commands
    return {'route': route, 'bound': len(route) if optimal else bound, 'optimal': optimal, 'lp': lp.value,
            'nodes': info.mip_node_count, 'seconds': time.perf_counter() - t0}
