#!/usr/bin/env python3

# Ein_System_Vier_Ansaetze.py
"""
Kapitel Oekosystem: Dasselbe LP in vier Bibliotheken.
   max 10*x1 + 15*x2 + 25*x3
   u.d.N. x1 + x2 + 2*x3 <= 40
          2*x1 + 3*x2 + x3 <= 50
          x >= 0

Deckt scipy.optimize, highspy, CVXPY und OR-Tools/GLOP ab, mit Kreuzvergleich
am Ende.

WICHTIG: Jeder Solver laeuft in einem EIGENEN Prozess, weil sich ortools und
highspy auf vielen Systemen nicht gemeinsam importieren lassen (beide bringen
eine eigene HiGHS-Kopie mit -> Symbolkonflikt).

Die Isolation besorgt ein ProcessPoolExecutor. Drei Einstellungen ergeben
zusammen die Garantie:

  mp_context "spawn"     Der Kindprozess startet mit einem FRISCHEN
                         Interpreter, statt den Speicher des Elternprozesses
                         zu erben. Was hier schon importiert ist, ist dort
                         nicht importiert. Mit dem Standard "fork" auf Linux
                         waere das nicht so.
  max_tasks_per_child=1  Jede Aufgabe bekommt einen NEUEN Prozess. Ohne das
                         wuerde der Pool seinen Arbeiter wiederverwenden - und
                         beim zweiten Solver waere der Konflikt zurueck.
  max_workers=1          Haelt die vier Laeufe nacheinander. Nicht aus
                         Vorsicht, sondern damit die gemessenen Zeiten
                         vergleichbar bleiben.

Jeder Solver steht in einer eigenen Funktion mit LOKALEM Import. Das ist der
Unterschied zu einem Codestring, den man an 'python -c' uebergibt: Die
Funktion laesst sich einzeln aufrufen, testen und vom Editor pruefen - ein
String nicht.

Benoetigt: scipy, highspy, cvxpy, ortools
"""

import multiprocessing
import time
from concurrent.futures import ProcessPoolExecutor

ERWARTET = 530.0            # Ergebnis der Handrechnung zum Produktionsprogramm

# Die Instanz - einmal notiert, von allen vier Funktionen benutzt.
ZIEL = [10.0, 15.0, 25.0]
MATRIX = [[1, 1, 2], [2, 3, 1]]
KAPAZITAET = [40.0, 50.0]


def loese_mit_scipy() -> tuple[float, list[float]]:
    from scipy.optimize import linprog
    ergebnis = linprog(c=[-w for w in ZIEL],          # linprog MINIMIERT -> negieren
                       A_ub=MATRIX, b_ub=KAPAZITAET,
                       bounds=[(0, None)] * 3, method="highs")
    return -ergebnis.fun, list(ergebnis.x)


def loese_mit_highspy() -> tuple[float, list[float]]:
    import highspy
    import numpy as np
    h = highspy.Highs()
    h.setOptionValue("output_flag", False)
    h.addVars(3, np.zeros(3), np.full(3, highspy.kHighsInf))
    h.changeObjectiveSense(highspy.ObjSense.kMaximize)
    for j, wert in enumerate(ZIEL):
        h.changeColCost(j, wert)
    # CSR-Format: starts[i] = Beginn von Zeile i in indices/values
    h.addRows(2, np.full(2, -highspy.kHighsInf), np.array(KAPAZITAET), 6,
              np.array([0, 3], dtype=np.int32),
              np.array([0, 1, 2, 0, 1, 2], dtype=np.int32),
              np.array([float(w) for zeile in MATRIX for w in zeile]))
    h.run()
    return (h.getInfo().objective_function_value,
            list(h.getSolution().col_value[:3]))


def loese_mit_cvxpy() -> tuple[float, list[float]]:
    import cvxpy as cp
    import numpy as np
    x = cp.Variable(3, nonneg=True)
    problem = cp.Problem(cp.Maximize(np.array(ZIEL) @ x),
                         [np.array(MATRIX) @ x <= np.array(KAPAZITAET)])
    problem.solve()
    return float(problem.value), [float(v) for v in x.value]


def loese_mit_ortools() -> tuple[float, list[float]]:
    from ortools.linear_solver import pywraplp
    s = pywraplp.Solver.CreateSolver("GLOP")
    x = [s.NumVar(0, s.infinity(), f"x{j+1}") for j in range(3)]
    for i, kapazitaet in enumerate(KAPAZITAET):
        s.Add(sum(MATRIX[i][j] * x[j] for j in range(3)) <= kapazitaet)
    s.Maximize(sum(ZIEL[j] * x[j] for j in range(3)))
    s.Solve()
    return s.Objective().Value(), [v.solution_value() for v in x]


ANSAETZE = {
    "scipy.optimize.linprog": loese_mit_scipy,
    "highspy (natives HiGHS)": loese_mit_highspy,
    "cvxpy": loese_mit_cvxpy,
    "ortools / GLOP": loese_mit_ortools,
}


if __name__ == "__main__":
    print("=" * 78)
    print("  EIN SYSTEM - VIER ANSAETZE (je eigener Prozess)")
    print("=" * 78)
    print(f"{'Bibliothek':<26} {'Z*':>10} {'x1':>7} {'x2':>7} {'x3':>7} {'Zeit':>10}")
    print("-" * 78)

    werte = []
    # Ein Pool, vier Aufgaben, vier frische Prozesse. Der Kontext muss
    # "spawn" sein - siehe Modulkommentar.
    with ProcessPoolExecutor(
            max_workers=1,
            mp_context=multiprocessing.get_context("spawn"),
            max_tasks_per_child=1) as pool:
        for name, funktion in ANSAETZE.items():
            beginn = time.perf_counter()
            try:
                wert, x = pool.submit(funktion).result(timeout=120)
            except Exception as fehler:                     # Bibliothek fehlt o. Ae.
                print(f"{name:<26} nicht verfuegbar: {str(fehler)[:40]}")
                continue
            dauer = time.perf_counter() - beginn
            werte.append(wert)
            print(f"{name:<26} {wert:>10.2f} {x[0]:>7.2f} {x[1]:>7.2f} {x[2]:>7.2f} "
                  f"{dauer:>8.2f} s")

    print("-" * 78)
    spanne = max(werte) - min(werte)
    print(f"Spannweite zwischen den Bibliotheken: {spanne:.2e}")
    print(f"Abweichung zur Handrechnung ({ERWARTET:.0f}):        "
          f"{abs(werte[0] - ERWARTET):.2e}")
    assert spanne < 1e-6, "Die Bibliotheken widersprechen sich!"
    assert abs(werte[0] - ERWARTET) < 1e-6, "Ergebnis weicht von der Handrechnung ab!"
    print("Alle Wege fuehren zum selben, von Hand bestaetigten Optimum.")
    print("(Die Zeiten enthalten Prozessstart und Import - sie messen NICHT die")
    print(" reine Solverleistung. Die Uebungsaufgabe 'Laufzeitvergleich' trennt beides.)")
    print("=" * 78)
