#!/usr/bin/env python3

# Simplex_Tableau_LP.py
"""
Kapitel LP: Vollständige Implementierung des Simplex-Algorithmus (Tableau-Methode).

GRENZEN DIESER IMPLEMENTIERUNG (bewusst, aus didaktischen Gründen):
  * nur Maximierung
  * nur "<="-Nebenbedingungen
  * alle b_i >= 0   (sonst wäre der Ursprung keine zulässige Startecke und man
    bräuchte eine Phase-1-Rechnung mit künstlichen Variablen)
Für den produktiven Einsatz nimmt man HiGHS - dieser Code dient dem Verständnis.

Voraussetzungen werden geprüft statt stillschweigend angenommen;
Iterationsprotokoll und Schattenpreise werden ausgegeben.
"""

import numpy as np


class SimplexTableauSolver:
    """Maximierungs-Standardform:  max c^T x  u.d.N.  A x <= b,  x >= 0,  b >= 0."""

    def __init__(self, c, A, b, variablennamen=None, restriktionsnamen=None):
        self.c = np.asarray(c, dtype=float)
        self.A = np.asarray(A, dtype=float)
        self.b = np.asarray(b, dtype=float)
        self.n = len(self.c)                    # Anzahl Originalvariablen
        self.m = len(self.b)                    # Anzahl Nebenbedingungen

        # --- Voraussetzungen pruefen, statt sie stillschweigend anzunehmen ---
        if self.A.shape != (self.m, self.n):
            raise ValueError(f"A hat Form {self.A.shape}, erwartet ({self.m}, {self.n}).")
        if np.any(self.b < 0):
            raise ValueError(
                "Mindestens ein b_i ist negativ. Dann ist der Ursprung keine zulässige "
                "Startecke; dieser Solver benötigt eine Phase-1-Rechnung, die hier "
                "bewusst nicht implementiert ist. Nutzen Sie scipy.optimize.linprog."
            )

        self.var_namen = variablennamen or [f"x{j+1}" for j in range(self.n)]
        self.restr_namen = restriktionsnamen or [f"R{i+1}" for i in range(self.m)]

        self.tableau = None
        self.basis = None                       # welche Variable ist in welcher Zeile Basis?
        self._baue_starttableau()

    def _baue_starttableau(self):
        """Zeilen: m Nebenbedingungen + Zielfunktionszeile.
           Spalten: n Variablen + m Schlupfvariablen + rechte Seite."""
        m, n = self.m, self.n
        self.tableau = np.zeros((m + 1, n + m + 1))
        self.tableau[:m, :n] = self.A                    # Koeffizienten
        self.tableau[:m, n:n + m] = np.eye(m)            # Schlupfvariablen
        self.tableau[:m, -1] = self.b                    # rechte Seite
        self.tableau[-1, :n] = -self.c                   # Zielzeile: -c (Maximierung)
        self.basis = list(range(n, n + m))               # Start: alle Schlupf in der Basis

    def _spaltenname(self, index):
        return self.var_namen[index] if index < self.n else f"s{index - self.n + 1}"

    def solve(self, max_iterationen=100, protokoll=True):
        m, n = self.m, self.n
        if protokoll:
            print(f"{'Iter':>4} | {'eintritt':>9} | {'austritt':>9} | "
                  f"{'Pivot':>8} | {'Z':>12}")
            print("-" * 58)

        for iteration in range(1, max_iterationen + 1):
            zielzeile = self.tableau[-1, :-1]

            # 1. Optimalitätsprüfung: alle Koeffizienten >= 0 ?
            if np.all(zielzeile >= -1e-9):
                if protokoll:
                    print("-" * 58)
                    print(f"Optimum nach {iteration - 1} Pivotschritten erreicht.")
                return self._loesung_auslesen()

            # 2. Pivotspalte: negativster Eintrag (Dantzig-Regel)
            pivot_spalte = int(np.argmin(zielzeile))

            # 3. Pivotzeile: minimaler Quotient über POSITIVE Spalteneinträge
            spalte = self.tableau[:m, pivot_spalte]
            rechte_seite = self.tableau[:m, -1]
            quotienten = np.where(spalte > 1e-9, rechte_seite / np.where(spalte > 1e-9, spalte, 1),
                                  np.inf)
            pivot_zeile = int(np.argmin(quotienten))
            if not np.isfinite(quotienten[pivot_zeile]):
                raise ValueError(
                    f"Problem ist unbeschraenkt: Variable {self._spaltenname(pivot_spalte)} "
                    "kann beliebig wachsen, ohne eine Bedingung zu verletzen. "
                    "Meist fehlt eine Kapazitaetsbeschraenkung."
                )

            if protokoll:
                print(f"{iteration:>4} | {self._spaltenname(pivot_spalte):>9} | "
                      f"{self._spaltenname(self.basis[pivot_zeile]):>9} | "
                      f"{self.tableau[pivot_zeile, pivot_spalte]:>8.3f} | "
                      f"{self.tableau[-1, -1]:>12,.2f}")

            # 4. Pivotoperation (Gauß-Jordan)
            pivot_wert = self.tableau[pivot_zeile, pivot_spalte]
            self.tableau[pivot_zeile, :] /= pivot_wert
            for zeile in range(m + 1):
                if zeile != pivot_zeile:
                    faktor = self.tableau[zeile, pivot_spalte]
                    self.tableau[zeile, :] -= faktor * self.tableau[pivot_zeile, :]

            self.basis[pivot_zeile] = pivot_spalte

        raise RuntimeError("Maximale Iterationszahl ueberschritten (moeglicherweise Zyklus).")

    def _loesung_auslesen(self):
        """Basisvariablen tragen den RHS-Wert ihrer Zeile, Nichtbasisvariablen sind 0."""
        x = np.zeros(self.n + self.m)
        for zeile, spalte in enumerate(self.basis):
            x[spalte] = self.tableau[zeile, -1]
        return x[:self.n], x[self.n:], self.tableau[-1, -1]

    def schattenpreise(self):
        """Die Zielzeile unter den Schlupfspalten enthält direkt die Dualwerte."""
        return self.tableau[-1, self.n:self.n + self.m].copy()


if __name__ == "__main__":
    # Modell aus der Simplex-Handrechnung (Bot-Beispiel, Kapitel Einfuehrung)
    ertraege = [150.0, 250.0]
    matrix = [[2.0, 5.0],
              [4.0, 6.0],
              [1.0, 0.0]]
    kapazitaeten = [40.0, 60.0, 8.0]
    var_namen = ["x_A", "x_B"]
    restr_namen = ["vCPU", "RAM", "Marktlimit"]

    print("=" * 58)
    print("  SIMPLEX-TABLEAU: ITERATIONSPROTOKOLL")
    print("=" * 58)

    solver = SimplexTableauSolver(ertraege, matrix, kapazitaeten, var_namen, restr_namen)
    x_opt, schlupf, z_opt = solver.solve()

    print("\n" + "=" * 58)
    print("  ERGEBNIS")
    print("=" * 58)
    for name, wert in zip(var_namen, x_opt):
        print(f"  {name:<12} = {wert:8.4f}")
    print(f"  {'Zielwert Z':<12} = {z_opt:8.2f} EUR")

    print("\n  Ressourcenanalyse:")
    print(f"  {'Ressource':<12} {'Schlupf':>9} {'Status':>22} {'Schattenpreis':>15}")
    print("  " + "-" * 60)
    for name, s, y in zip(restr_namen, schlupf, solver.schattenpreise()):
        status = "ENGPASS (bindend)" if abs(s) < 1e-9 else "Reserve vorhanden"
        print(f"  {name:<12} {s:>9.3f} {status:>22} {y:>12.2f} EUR")

    # Selbstkontrolle: komplementaerer Schlupf muss gelten
    for s, y in zip(schlupf, solver.schattenpreise()):
        assert abs(s * y) < 1e-6, "Komplementaerer Schlupf verletzt - Rechenfehler!"
    print("\n  Pruefung: komplementaerer Schlupf (s_i * y_i = 0) fuer alle i erfuellt.")
    print("=" * 58)
