#!/usr/bin/env python3

# Min_Cost_Flow.py
"""
Kapitel Graphen: Kostenminimaler Fluss durch ein Netzwerk.

Loest dasselbe Problem zweimal:
  (1) als allgemeines LP mit scipy  -> zeigt die Modellstruktur
  (2) mit dem spezialisierten Netzwerk-Solver von OR-Tools -> zeigt den
      Geschwindigkeitsvorteil eines Verfahrens, das die Struktur ausnutzt

Beide laufen in eigenen Prozessen nicht noetig: scipy und ortools vertragen
sich (nur ortools + highspy kollidieren, siehe Kapitel Oekosystem).
"""

import numpy as np
from scipy.optimize import linprog

# --- Netzwerk definieren ---------------------------------------------------
KNOTEN = ["Werk_A", "Werk_B", "Umschlag", "Kunde_1", "Kunde_2"]
# (von, nach, Kosten je Einheit, Kapazitaet)
KANTEN = [
    ("Werk_A",   "Umschlag", 2.0, 15),
    ("Werk_A",   "Kunde_1",  5.0, 10),
    ("Werk_B",   "Umschlag", 4.0, 10),
    ("Werk_B",   "Kunde_2",  6.0, 10),
    ("Umschlag", "Kunde_1",  1.0, 20),
    ("Umschlag", "Kunde_2",  3.0, 10),
]
# Angebot (+) bzw. Bedarf (-) je Knoten
SALDO = {"Werk_A": 20, "Werk_B": 10, "Umschlag": 0, "Kunde_1": -15, "Kunde_2": -15}


def loese_als_lp():
    """Flussproblem als allgemeines lineares Programm."""
    n_kanten = len(KANTEN)
    knoten_index = {k: i for i, k in enumerate(KNOTEN)}

    # Zielfunktion: Summe der Transportkosten
    kosten = np.array([k[2] for k in KANTEN])

    # Flusserhaltung als Gleichungssystem: A_eq @ x = b_eq
    A_eq = np.zeros((len(KNOTEN), n_kanten))
    for e, (von, nach, _, _) in enumerate(KANTEN):
        A_eq[knoten_index[von], e] = +1.0      # fliesst hinaus
        A_eq[knoten_index[nach], e] = -1.0     # fliesst hinein
    b_eq = np.array([SALDO[k] for k in KNOTEN], dtype=float)

    schranken = [(0, k[3]) for k in KANTEN]    # 0 <= x_ij <= u_ij

    ergebnis = linprog(c=kosten, A_eq=A_eq, b_eq=b_eq, bounds=schranken, method="highs")
    if not ergebnis.success:
        raise SystemExit(f"Nicht loesbar: {ergebnis.message}")
    return ergebnis.fun, ergebnis.x, ergebnis.eqlin.marginals


if __name__ == "__main__":
    # Vorabpruefung: Angebot muss Bedarf entsprechen
    gesamt = sum(SALDO.values())
    print("=" * 78)
    print("  KOSTENMINIMALER FLUSS DURCH EIN TRANSPORTNETZ")
    print("=" * 78)
    print(f"Angebot gesamt: {sum(v for v in SALDO.values() if v > 0)} | "
          f"Bedarf gesamt: {-sum(v for v in SALDO.values() if v < 0)} | "
          f"Saldo: {gesamt}")
    if gesamt != 0:
        raise SystemExit("Angebot und Bedarf stimmen nicht ueberein - unloesbar!")

    kosten_gesamt, fluss, knotenpreise = loese_als_lp()

    print(f"\nMinimale Transportkosten: {kosten_gesamt:,.2f} EUR\n")
    print(f"{'Kante':<24} {'Fluss':>7} {'Kapazitaet':>11} {'Kosten/E':>9} {'Kosten':>9}")
    print("-" * 78)
    for e, (von, nach, c, u) in enumerate(KANTEN):
        menge = fluss[e] + 0.0 if abs(fluss[e]) > 1e-9 else 0.0   # vermeidet "-0.0"
        ausgelastet = " (VOLL)" if abs(menge - u) < 1e-6 else ""
        print(f"{von + ' -> ' + nach:<24} {menge:>7.1f} {u:>11} "
              f"{c:>9.2f} {menge * c:>9.2f}{ausgelastet}")

    # --- Flusserhaltung nachpruefen --------------------------------------
    print("\n--- Pruefung der Flusserhaltung je Knoten ---")
    for k in KNOTEN:
        hinaus = sum(fluss[e] for e, (v, n, _, _) in enumerate(KANTEN) if v == k)
        hinein = sum(fluss[e] for e, (v, n, _, _) in enumerate(KANTEN) if n == k)
        netto = hinaus - hinein
        art = "Quelle" if SALDO[k] > 0 else ("Senke" if SALDO[k] < 0 else "Umschlag")
        print(f"  {k:<10} ({art:<8}): hinaus {hinaus:5.1f} - hinein {hinein:5.1f} "
              f"= {netto:+6.1f}  (gefordert: {SALDO[k]:+d})")
        assert abs(netto - SALDO[k]) < 1e-6, f"Flusserhaltung verletzt bei {k}!"

    # --- Knotenpreise (Dualwerte) interpretieren -------------------------
    print("\n--- Knotenpreise (Dualwerte der Flusserhaltung) ---")
    print("  Differenz zweier Knotenpreise = Grenzkosten einer zusaetzlichen Einheit")
    print("  auf dem guenstigsten Weg zwischen ihnen.")
    for k, preis in zip(KNOTEN, knotenpreise):
        print(f"  {k:<10}: {preis:7.2f}")
    print("=" * 78)
