"""Part 8: how short can the robot's route be? A lower bound from a linear program.

Part 3 asked for the fewest commands that clean the whole floor (Part 1's rules, finish
anywhere). A route is a walk from the start s. Close it with a free edge from wherever it ends
back to s: every field then has even degree, and the walk is connected. So:

    x_e in {0, 1, 2}   times the route walks between neighbouring fields e = {u, v}   (cost 1)
    z_v in {0, 1}      the route ends on v; exactly one z_v is 1                    (cost 0)
    parity             x(delta(v)) + z_v is even for v != s, x(delta(s)) + 1 - z_s is even
    connected          x(delta(S)) + z(S) >= 2 for every non-empty set S of fields without s

(No route needs an edge three times: drop two copies and parity and connectivity survive.)
Every route is a solution, so the minimum of sum x_e over the solutions is at most the shortest
route, and in fact equal to it. Drop parity and let x and z be fractions and it becomes a linear
program, whose minimum is a lower bound on every route. Part 9 puts parity back.

There is one "connected" row for every set S, far too many to write down, so the program starts
with S = {v} for each field and every 2 x 2 block, and adds a row only when the current answer
breaks it: the cutting
plane method. Finding a broken row is a minimum cut. Give each edge e the capacity x_e, and the
closing edge from v to s the capacity z_v. A set S with x(delta(S)) + z(S) < 2 is then exactly a
cut of capacity under 2 separating s from some field.

Two solvers do the same loop: simplex.solve, exact and from scratch, for small rooms, and HiGHS
(the highspy package) for big ones.
"""
import math
from fractions import Fraction

import simplex
from maxflow import Network
from shortest_route_fast import parse

LOOSE = 1e-7            # HiGHS works in floating point: violations smaller than this don't count


class Model:
    """The room as a graph: fields 0..n-1 (sorted), edges between neighbours, the start s."""

    def __init__(self, room):
        start, floor = parse(room)
        self.fields = sorted(floor)
        index = {f: i for i, f in enumerate(self.fields)}
        self.s = index[start]
        self.edges = [(index[(r, c)], index[(r + dr, c + dc)])
                      for r, c in self.fields for dr, dc in ((1, 0), (0, 1)) if (r + dr, c + dc) in index]
        self.n, self.m = len(self.fields), len(self.edges)
        self.touching = [[] for _ in range(self.n)]
        for k, (a, b) in enumerate(self.edges):
            self.touching[a].append(k)
            self.touching[b].append(k)
        self.rows = {}

    def first_cuts(self):
        """The rows the loop starts with: every field on its own, and every 2 x 2 block of fields.
        A loop round a 2 x 2 block, cut off from the rest, is the commonest way the program cheats,
        so ruling it out from the start saves many rounds."""
        index = {f: i for i, f in enumerate(self.fields)}
        cuts = [frozenset([v]) for v in range(self.n) if v != self.s]
        for r, c in self.fields:
            block = [index.get(f) for f in ((r, c), (r, c + 1), (r + 1, c), (r + 1, c + 1))]
            if None not in block and self.s not in block:
                cuts.append(frozenset(block))
        return cuts

    def cut_row(self, S):
        """The row x(delta(S)) + z(S) >= 2, as {variable: coefficient}: x_e is variable e, z_v is
        variable m + v."""
        if S in self.rows:
            return self.rows[S]
        row = self.rows[S] = {}
        for v in S:
            row[self.m + v] = 1
            for k in self.touching[v]:
                a, b = self.edges[k]
                if (a in S) != (b in S):
                    row[k] = 1
        return row

    def crossing(self, S, x, z):
        return sum(z[v] for v in S) + sum(x[k] for v in S for k in self.touching[v]
                                          if (self.edges[k][0] in S) != (self.edges[k][1] in S))


def broken_cuts(model, x, z, eps=0):
    """Sets S (without s) whose row x(delta(S)) + z(S) >= 2 the answer x, z breaks.

    Cheap tries first: the connected pieces of the graph of edges carrying at least theta, for a
    few thresholds theta. Only when those find nothing, a minimum cut from s to every field."""
    n, s = model.n, model.s
    edges = [(a, b, x[k]) for k, (a, b) in enumerate(model.edges) if x[k] > eps]
    edges += [(v, s, z[v]) for v in range(n) if v != s and z[v] > eps]
    found = []
    seen = set()
    steps = (Fraction(1, 4), Fraction(1, 2), Fraction(3, 4), Fraction(99, 100))
    for theta in (eps,) + (steps if eps == 0 else tuple(float(t) for t in steps)):
        for S in pieces(n, [(a, b) for a, b, w in edges if w > theta], s):
            if S not in seen and model.crossing(S, x, z) < 2 - eps:
                seen.add(S)
                found.append(S)
    if found:
        return found
    network = Network(n, edges)
    for v in range(n):
        if v == s:
            continue
        value, side = network.min_cut(s, v, enough=2, eps=eps)
        if side is not None and value < 2 - eps:
            S = frozenset(range(n)) - side
            if S not in seen:
                seen.add(S)
                found.append(S)
    return found


def pieces(n, pairs, s):
    """The connected pieces of the graph on n nodes with these edges, except the one holding s."""
    adj = [[] for _ in range(n)]
    for a, b in pairs:
        adj[a].append(b)
        adj[b].append(a)
    label = [-1] * n
    out = []
    for i in range(n):
        if label[i] >= 0:
            continue
        label[i] = i
        stack, piece = [i], [i]
        while stack:
            u = stack.pop()
            for v in adj[u]:
                if label[v] < 0:
                    label[v] = i
                    stack.append(v)
                    piece.append(v)
        if label[s] != i:
            out.append(frozenset(piece))
    return out


class Bound:
    """What the cutting-plane loop found."""

    def __init__(self, model, value, x, z, cuts, rounds, duals, history):
        self.model = model
        self.value = value              # the LP's minimum (Fraction with the exact solver)
        self.commands = math.ceil(value - LOOSE)    # a route has a whole number of commands
        self.x, self.z = x, z           # the optimal fractional solution
        self.cuts = cuts                # every S whose row the program ended with
        self.rounds = rounds
        self.duals = duals              # price of each cut row, then of sum z = 1
        self.history = history          # [(value, cuts added), ...], one per round


def lp_bound(room, solver='exact', max_rounds=10_000, on_round=None):
    """Run the cutting-plane loop. solver: 'exact' (simplex.py, from scratch) or 'highs' (HiGHS,
    through its own Python package highspy, which keeps the program between rounds and restarts the
    dual simplex from the last basis). on_round(value, x, z, broken) is called after every round."""
    model = Model(room)
    if model.n == 1:
        return Bound(model, 0, [], [0], [], 0, [], [])
    lp = ExactLP(model) if solver == 'exact' else HighsLP(model)
    eps = 0 if solver == 'exact' else LOOSE
    new, history = model.first_cuts(), []
    for rounds in range(1, max_rounds + 1):
        lp.add(new)
        value, x, z, duals = lp.solve()
        new = broken_cuts(model, x, z, eps)
        history.append((value, len(new)))
        if on_round:
            on_round(value, x, z, new)
        if not new:
            return Bound(model, value, x, z, lp.cuts, rounds, duals, history)
    raise RuntimeError('the cutting-plane loop did not settle')


def rows_for(model, cuts):
    n, m = model.n, model.m
    rows = [(model.cut_row(S), simplex.GE, 2) for S in cuts]
    rows.append(({m + v: 1 for v in range(n)}, simplex.EQ, 1))
    return rows


class ExactLP:
    """The program in exact fractions, solved from scratch by simplex.py every round."""

    def __init__(self, model):
        self.model, self.cuts = model, []

    def add(self, cuts):
        self.cuts += cuts

    def solve(self):
        n, m = self.model.n, self.model.m
        rows = rows_for(self.model, self.cuts)
        rows += [({k: 1}, simplex.LE, 2) for k in range(m)]
        rows += [({m + v: 1}, simplex.LE, 1) for v in range(n)]
        result = simplex.solve([1] * m + [0] * n, rows)
        assert result.status == 'optimal'
        return result.value, result.x[:m], result.x[m:], result.y[:len(self.cuts) + 1]


class HighsLP:
    """The same program in HiGHS. Row 0 is sum z = 1; cut rows follow in the order they came."""

    def __init__(self, model):
        import highspy
        import numpy as np
        self.np, self.model, self.cuts = np, model, []
        n, m = model.n, model.m
        h = self.h = highspy.Highs()
        h.setOptionValue('output_flag', False)
        h.addVars(m + n, np.zeros(m + n), np.r_[np.full(m, 2.0), np.ones(n)])
        h.changeColsCost(m, np.arange(m, dtype=np.int32), np.ones(m))
        h.addRow(1.0, 1.0, n, np.arange(m, m + n, dtype=np.int32), np.ones(n))
        self.optimal = highspy.HighsModelStatus.kOptimal

    def add(self, cuts):
        np = self.np
        starts, index, values = [], [], []
        for S in cuts:
            row = self.model.cut_row(S)
            starts.append(len(index))
            index += row.keys()
            values += row.values()
        self.h.addRows(len(cuts), np.full(len(cuts), 2.0), np.full(len(cuts), np.inf), len(index),
                       np.array(starts, dtype=np.int32), np.array(index, dtype=np.int32),
                       np.array(values, dtype=float))
        self.cuts += cuts

    def solve(self):
        h, m = self.h, self.model.m
        h.run()
        assert h.getModelStatus() == self.optimal, h.modelStatusToString(h.getModelStatus())
        sol = h.getSolution()
        col, dual = list(sol.col_value), list(sol.row_dual)
        return h.getInfo().objective_function_value, col[:m], col[m:], dual[1:] + dual[:1]
