"""Part 9: the fewest commands, exactly. Part 8's program with parity and whole numbers put back.

Part 8 dropped two kinds of rule to get a linear program: every x_e, z_v a whole number, and
every field touched by an even number of steps. Put back, they make an integer program whose
minimum is exactly the shortest route. Parity is a whole-number rule in disguise:

    x(delta(v)) + z_v - 2 w_v = 0          for every field v other than the start,
    x(delta(s)) - z_s - 2 w_s = -1         for the start (the closing step always touches it),

with w_v a whole number, which says the left sides are even. Linear programs can't say "whole
number", so branch and bound says it for them:

- Solve the linear program (with Part 8's cutting-plane loop for the connected rules). Its value
  is a lower bound for every route in this part of the search.
- If that bound can't beat the best route found so far, stop here: no route below can either.
- If the answer is all whole numbers, it is a route: keep it if it is the best so far.
- Otherwise pick a variable with a fractional value v and split: one branch adds "<= floor(v)",
  the other ">= ceil(v)". Every whole-number answer is in one of them, and the fractional one is
  in neither.

Gomory's cuts, from the final tableau, can be added at the root first: rules every whole-number
answer obeys and the fractional answer breaks. They make the tree smaller.

Everything here is exact (simplex.py, fractions), so it is only for small rooms. route_ip.py
solves the same program with HiGHS on big ones.
"""
import math
from fractions import Fraction

import route_lp
import simplex
from shortest_route_fast import solution as fast_route


class Program:
    """The integer program for one room. Variables: x_e (m of them), then z_v (n), then w_v (n)."""

    def __init__(self, room):
        self.model = model = route_lp.Model(room)
        n, m, s = model.n, model.m, model.s
        self.size = m + 2 * n
        self.cost = [1] * m + [0] * (2 * n)
        self.base = [({m + v: 1 for v in range(n)}, simplex.EQ, 1)]
        for v in range(n):
            row = {k: 1 for k in model.touching[v]}
            row[m + v] = -1 if v == s else 1
            row[m + n + v] = -2
            self.base.append((row, simplex.EQ, -1 if v == s else 0))
        self.upper = [2] * m + [1] * n + [4] * n        # a field has at most 4 neighbours: degree <= 8
        self.cuts = model.first_cuts()                  # the connected rules found so far, shared
        self.gomory = []                                # rows from Gomory's cuts at the root
        self.lps = 0

    def solve(self, bounds):
        """The linear program with extra bounds {variable: (low, high)}, and the cutting-plane loop
        for the connected rules. Returns (value, values) or None if infeasible."""
        m, n = self.model.m, self.model.n
        while True:
            rows = list(self.base) + [(self.model.cut_row(S), simplex.GE, 2) for S in self.cuts] + list(self.gomory)
            for j in range(self.size):
                low, high = bounds.get(j, (0, self.upper[j]))
                rows.append(({j: 1}, simplex.LE, high))
                if low:
                    rows.append(({j: 1}, simplex.GE, low))
            result = simplex.solve(self.cost, rows)
            self.lps += 1
            if result.status != 'optimal':
                return None
            x, z = result.x[:m], result.x[m:m + n]
            broken = route_lp.broken_cuts(self.model, x, z)
            if not broken:
                self.last = result
                return result.value, result.x
            self.cuts += broken


def fractional(values):
    """The variable to branch on: the one whose value is furthest from a whole number."""
    best, pick = Fraction(0), None
    for j, v in enumerate(values):
        f = v - math.floor(v)
        d = min(f, 1 - f)
        if d > best:
            best, pick = d, j
    return pick


def solve(room, gomory_rounds=0, record=None):
    """The shortest route, by branch and bound. Returns (commands, stats).

    record, if given, is called for every node: record(node, parent, branch, bound, outcome, x, z),
    with the node's LP answer (None where the LP is infeasible)."""
    p = Program(room)
    m, n = p.model.m, p.model.n
    best = fast_route(room)                             # Part 3's fast route: the first incumbent
    stats = {'nodes': 0, 'gomory': 0, 'root': None}
    # Gomory's cuts at the root: valid for every whole-number answer, so for the whole tree.
    for _ in range(gomory_rounds):
        solved = p.solve({})
        if solved is None or fractional(solved[1]) is None:
            break
        cuts = simplex.gomory_cuts(p.last)
        if not cuts:
            break
        p.gomory += cuts
        stats['gomory'] += len(cuts)
    stack = [({}, None, None)]
    while stack:
        bounds, parent, branch = stack.pop()
        node = stats['nodes']
        stats['nodes'] += 1
        solved = p.solve(bounds)
        x = z = None
        if solved is None:
            outcome, bound = 'infeasible', None
        else:
            value, values = solved
            bound = value
            x, z = values[:m], values[m:m + n]
            if parent is None:
                stats['root'] = value
            if math.ceil(value) >= len(best):
                outcome = 'pruned'
            else:
                j = fractional(values)
                if j is None:
                    route = euler_route(p.model, [int(v) for v in values[:m]], [int(v) for v in values[m:m + n]])
                    assert len(route) == value
                    best, outcome = route, 'route'
                else:
                    outcome = 'branched'
                    low, high = bounds.get(j, (0, p.upper[j]))
                    below, above = math.floor(values[j]), math.ceil(values[j])
                    down_bounds, up_bounds = dict(bounds), dict(bounds)
                    down_bounds[j], up_bounds[j] = (low, below), (above, high)
                    down = (down_bounds, node, (j, '<=', below))
                    up = (up_bounds, node, (j, '>=', above))
                    # the side nearer the fractional value is searched first, so it goes on last
                    near_up = values[j] - below >= Fraction(1, 2)
                    stack += [down, up] if near_up else [up, down]
        if record:
            record(node, parent, branch, bound, outcome, x, z)
    stats['lps'] = p.lps
    stats['cuts'] = len(p.cuts)
    return best, stats


def euler_route(model, x, z):
    """Turn whole-number x (walks per pair of neighbours) and z (where it ends) into commands.

    Every field has even degree except the start and the end (Part 9's parity rules), and the
    steps are connected (the cut rules), so Hierholzer's algorithm walks them all in one go:
    follow unused steps until stuck, then back up and splice in the loops left behind."""
    s = model.s
    t = z.index(1)
    left = [dict() for _ in range(model.n)]
    for k, (a, b) in enumerate(model.edges):
        if x[k]:
            left[a][b] = left[a].get(b, 0) + x[k]
            left[b][a] = left[b].get(a, 0) + x[k]
    stack, path = [s], []
    while stack:
        u = stack[-1]
        if left[u]:
            v = next(iter(left[u]))
            for a, b in ((u, v), (v, u)):
                left[a][b] -= 1
                if not left[a][b]:
                    del left[a][b]
            stack.append(v)
        else:
            path.append(stack.pop())
    path.reverse()
    assert path[0] == s and path[-1] == t
    go = {(-1, 0): '^', (1, 0): 'v', (0, -1): '<', (0, 1): '>'}
    f = model.fields
    return ''.join(go[(f[b][0] - f[a][0], f[b][1] - f[a][1])] for a, b in zip(path, path[1:]))
