"""Checks and a report for Part 8: a lower bound on the fewest commands, from a linear program.

Run from this folder: python test_route_lp.py
The HiGHS checks and the big-room report need highspy (pip install highspy); without it they are skipped, and say so.
"""
import itertools
import random
import time
from fractions import Fraction

import route_lp
import shortest_route_exact
import simplex
from cleaning_robot import solution_trimmed as part1_trimmed
from maxflow import min_cut
from shortest_route_fast import lower_bound, parse
from shortest_route_fast import solution as fast
from test_cleaning_robot import EXAMPLES, random_room
from test_shortest_route import CORRIDOR, TWO_LOOPS

try:
    import highspy  # noqa: F401
    HAVE_HIGHS = True
except ImportError:
    HAVE_HIGHS = False

EXAMPLE_2 = EXAMPLES[1][0]          # Part 3's example 2: 5x5, start in the middle


def random_lp(rng):
    n, m = rng.randint(1, 5), rng.randint(1, 5)
    c = [rng.randint(-4, 5) for _ in range(n)]
    rows = []
    for _ in range(m):
        coef = {j: rng.randint(-3, 4) for j in range(n) if rng.random() < 0.8}
        rows.append((coef, rng.choice([simplex.LE, simplex.GE, simplex.EQ]), rng.randint(-5, 8)))
    return c, rows


def highs_lp(c, rows):
    """The same program in HiGHS, with presolve off so it says which of the three it is."""
    import highspy
    import numpy as np
    n = len(c)
    h = highspy.Highs()
    h.setOptionValue('output_flag', False)
    h.setOptionValue('presolve', 'off')
    h.addVars(n, np.zeros(n), np.full(n, np.inf))
    h.changeColsCost(n, np.arange(n, dtype=np.int32), np.array(c, dtype=float))
    for coef, op, b in rows:
        lo = b if op in (simplex.GE, simplex.EQ) else -np.inf
        hi = b if op in (simplex.LE, simplex.EQ) else np.inf
        h.addRow(lo, hi, len(coef), np.array(list(coef), dtype=np.int32), np.array(list(coef.values()), dtype=float))
    h.run()
    status = {highspy.HighsModelStatus.kOptimal: 'optimal', highspy.HighsModelStatus.kInfeasible: 'infeasible',
              highspy.HighsModelStatus.kUnbounded: 'unbounded'}[h.getModelStatus()]
    return status, h.getInfo().objective_function_value


def test_simplex(rng, count=600):
    """The exact simplex against HiGHS on random programs, and every optimum's certificate."""
    kinds = {'optimal': 0, 'infeasible': 0, 'unbounded': 0}
    for _ in range(count):
        c, rows = random_lp(rng)
        r = simplex.solve(c, rows)
        kinds[r.status] += 1
        if r.status == 'optimal':
            assert simplex.check_certificate(c, rows, r.x, r.y)
        if HAVE_HIGHS:
            status, value = highs_lp(c, rows)
            assert status == r.status
            if status == 'optimal':
                assert abs(float(r.value) - value) < 1e-7
    against = 'the same answer as HiGHS on every one' if HAVE_HIGHS else 'HiGHS skipped (no highspy)'
    print(f'simplex ................. ok ({count} random programs: {kinds["optimal"]} optimal, '
          f'{kinds["infeasible"]} infeasible, {kinds["unbounded"]} unbounded; {against}; '
          f'every optimum certified by its duals)')


def test_min_cut(rng, count=300):
    """Edmonds-Karp against every possible cut, on small random graphs."""
    for _ in range(count):
        n = rng.randint(2, 8)
        edges = [(a, b, Fraction(rng.randint(1, 6), rng.randint(1, 3)))
                 for a, b in itertools.combinations(range(n), 2) if rng.random() < 0.5]
        s, t = rng.sample(range(n), 2)
        value, side = min_cut(n, edges, s, t)
        assert s in side and t not in side
        assert value == sum(w for a, b, w in edges if (a in side) != (b in side))
        best = min(sum(w for a, b, w in edges if (a in S) != (b in S))
                   for k in range(n) for S in map(set, itertools.combinations(range(n), k)) if s in S and t not in S)
        assert value == best
    print(f'minimum cut ............. ok ({count} random graphs: the flow equals the smallest of every cut)')


def test_examples():
    for room, want in ((CORRIDOR, 7), (TWO_LOOPS, 7), (EXAMPLE_2, 24)):
        b = route_lp.lp_bound(room)
        assert b.commands == want == len(shortest_route_exact.solution(room)), (room, b.value)
    print('Part 3\'s examples ....... ok (the corridor 7, the two loops 7, example 2 24: the bound is the answer)')


def certify(bound):
    """The final LP's duals prove its minimum, and no cut row is broken: then bound.value is the
    LP's true minimum over all its (exponentially many) rows."""
    model = bound.model
    rows = route_lp.rows_for(model, bound.cuts)
    rows += [({k: 1}, simplex.LE, 2) for k in range(model.m)]
    rows += [({model.m + v: 1}, simplex.LE, 1) for v in range(model.n)]
    result = simplex.solve([1] * model.m + [0] * model.n, rows)
    assert result.value == bound.value
    assert simplex.check_certificate([1] * model.m + [0] * model.n, rows, result.x, result.y)
    assert not route_lp.broken_cuts(model, bound.x, bound.z)


def test_small_rooms(rng, count=300):
    """On rooms small enough for Part 3's exact search: the bound never exceeds the true answer,
    is never below Part 3's bound, and the exact and HiGHS loops agree."""
    tight = above = 0
    t_exact = t_highs = 0.0
    for _ in range(count):
        room = random_room(rng, rng.randint(3, 6), rng.randint(3, 6), rng.choice([0, 0.2, 0.35]))
        t = time.perf_counter()
        b = route_lp.lp_bound(room)
        t_exact += time.perf_counter() - t
        certify(b)
        best = len(shortest_route_exact.solution(room))
        assert lower_bound(room) <= b.commands <= best
        tight += b.commands == best
        above += b.commands > lower_bound(room)
        if HAVE_HIGHS:
            t = time.perf_counter()
            h = route_lp.lp_bound(room, 'highs')
            t_highs += time.perf_counter() - t
            assert abs(h.value - float(b.value)) < 1e-6
    extra = (f'; HiGHS agreed on all, {1000 * t_highs / count:.0f} ms against '
             f'{1000 * t_exact / count:.0f} ms exact') if HAVE_HIGHS else '; HiGHS skipped (no highspy)'
    print(f'small rooms ............. ok ({count} rooms up to 4x4 inside: never above the best route, '
          f'never below Part 3\'s bound; equal to the best on {tight}, above Part 3\'s bound on {above}; '
          f'every LP certified{extra})')


def report_duals(room, name):
    """The certificate: the price on each cut row. Every field with a price is a 'moat' the route
    must cross into and out of; 2 x (sum of prices) + (price of sum z = 1) is the bound."""
    b = route_lp.lp_bound(room)
    fields = b.model.fields
    priced = [(y, sorted(fields[v] for v in S)) for S, y in zip(b.cuts, b.duals) if y]
    singles = [(y, S) for y, S in priced if len(S) == 1]
    bigger = [(y, S) for y, S in priced if len(S) > 1]
    print(f'\n{name}: bound {b.value} = 2 x {sum(y for y, _ in priced)} (prices on {len(priced)} cut rows) '
          f'{b.duals[-1]:+} (the free end)')
    print(f'  {len(singles)} single fields at price {sorted({str(y) for y, _ in singles})}; '
          f'{len(bigger)} bigger sets: {[(str(y), len(S)) for y, S in bigger]}')


def report_rooms(rng, groups=((12, 0.15, 20), (12, 0.3, 20), (20, 0.15, 10), (20, 0.3, 10),
                              (40, 0.15, 3), (40, 0.3, 3))):
    """Random furnished rooms: how much of the gap between Part 3's fast route and Part 3's bound
    the LP bound closes, and how many fast routes it proves best."""
    print('\nRandom furnished rooms (HiGHS):')
    print('  size   furniture  rooms  fields  proven by Part 3  proven by LP  gap to Part 3  gap to LP  '
          'closed  rounds  time')
    for size, p, count in groups:
        rows = []
        for _ in range(count):
            room = random_room(rng, size, size, p)
            f, lb = len(fast(room)), lower_bound(room)
            t = time.perf_counter()
            b = route_lp.lp_bound(room, 'highs')
            rows.append((b.model.n, f, lb, b.commands, b.rounds, time.perf_counter() - t))
        n = sum(r[0] for r in rows) / count
        gap3 = sum(r[1] - r[2] for r in rows) / count
        gap8 = sum(r[1] - r[3] for r in rows) / count
        print(f'  {size}x{size}  {100 * p:>6.0f}%  {count:>6}  {n:>6.0f}  {sum(r[1] == r[2] for r in rows):>16}  '
              f'{sum(r[1] == r[3] for r in rows):>12}  {gap3:>13.1f}  {gap8:>9.1f}  '
              f'{100 * (1 - gap8 / gap3) if gap3 else 100:>5.0f}%  {sum(r[4] for r in rows) / count:>6.0f}  '
              f'{sum(r[5] for r in rows) / count:>4.1f}s')


def report_showcase():
    """Three 40x40 rooms, as in Part 3's table, and the room the figures use."""
    rng = random.Random(8)
    print('\nOne room of each kind (random.Random(8)):')
    print('  room          fields  Part 1 trimmed  fast  Part 3 bound  LP bound  rounds  cuts   time')
    for size, p in ((12, 0.2), (20, 0.3), (40, 0.3)):
        room = random_room(rng, size, size, p)
        t = time.perf_counter()
        b = route_lp.lp_bound(room, 'highs')
        print(f'  {size}x{size}, {100 * p:.0f}%  {b.model.n:>6}  {len(part1_trimmed(room)):>14}  {len(fast(room)):>4}  '
              f'{lower_bound(room):>12}  {b.value:>8.2f}  {b.rounds:>6}  {len(b.cuts):>4}  {time.perf_counter() - t:>5.1f}s')


def main():
    rng = random.Random(8)
    test_simplex(rng)
    test_min_cut(rng)
    test_examples()
    test_small_rooms(rng)
    report_duals(CORRIDOR, 'The corridor')
    report_duals(EXAMPLE_2, 'Example 2')
    if HAVE_HIGHS:
        report_showcase()
        report_rooms(rng)
    else:
        print('\nThe big-room reports need highspy (HiGHS); skipped.')


if __name__ == '__main__':
    main()
