#!/usr/bin/env python3

# Solverwechsel_CPSAT_HiGHS.py
"""
Kapitel Praxisfallen: Denselben Fall einmal mit CP-SAT und einmal mit HiGHS rechnen.

Die Trennung aus or_kern.py behauptet, ein Solverwechsel koste genau EINEN
Baustein. Dieses Programm loest diese Behauptung ein. Aufgabe ist eine
Standortplanung: Welche Lager oeffnen wir, und wer beliefert welchen Kunden?

    Standortproblem  (Pydantic, geprueft)   <- gemeinsam
        baue_und_loese_mit_cpsat()          <- der EINE Baustein
        baue_und_loese_mit_highs()          <- ... in zwei Ausfuehrungen
    Loesung          (DTO)                  <- gemeinsam
    pruefe_zuordnung()                      <- gemeinsam
    berichte()                              <- gemeinsam

Die beiden Modellbauer sind 39 und 52 Zeilen lang - sie sind der gesamte
solverabhaengige Teil des Programms. Alles andere wird zweimal benutzt und
einmal geschrieben.

WARUM ZWEI PROZESSE? ortools und highspy bringen beide eine eigene
HiGHS-Kopie mit und lassen sich auf vielen Systemen nicht gemeinsam
importieren (Kapitel Oekosystem). Das Hauptprogramm startet deshalb fuer jeden
Solver einen eigenen Python-Prozess und laesst sich die Loesung als JSON
zurueckgeben - das DTO ist nicht nur eine Sprachregelung, sondern ein
Datenformat, das eine Prozessgrenze ueberlebt.

Aufruf:
    python3 Solverwechsel_CPSAT_HiGHS.py            # beide, mit Vergleich
    python3 Solverwechsel_CPSAT_HiGHS.py cpsat      # nur der Kindprozess
    python3 Solverwechsel_CPSAT_HiGHS.py highs

Benoetigt: numpy, pydantic, ortools, highspy (jeweils im eigenen Prozess)
"""

from __future__ import annotations

import multiprocessing
import time
from concurrent.futures import ProcessPoolExecutor

import numpy as np
from pydantic import BaseModel, Field, model_validator

from or_kern import Loesung, SolverStatus, status_von_cpsat, status_von_highs

ZEITLIMIT = 30.0


# --- 1. Domaenenmodell: dieselben Fakten fuer beide Solver -------------------

class Standortproblem(BaseModel):
    """Kapazitierte Standortplanung mit Einzelbelieferung.

    Nach dem Muster von Produktionsproblem in or_kern.py: Die Pruefungen
    stehen im Konstruktor, nicht im Solvercode - sie gelten damit fuer beide
    Solver, und sie schlagen beim Einlesen zu.

    Alle Kosten sind ganzzahlig (Euro). Das ist keine Bequemlichkeit, sondern
    Voraussetzung: CP-SAT rechnet ausschliesslich ganzzahlig.
    """
    lager: list[str] = Field(min_length=1)
    kunden: list[str] = Field(min_length=1)
    fixkosten: list[int]        # je Lager, faellt bei Eroeffnung an
    kapazitaet: list[int]       # je Lager, in Paletten
    bedarf: list[int]           # je Kunde, in Paletten
    transport: list[list[int]]  # [Lager][Kunde], Kosten der Belieferung

    @model_validator(mode="after")
    def pruefe_masse(self) -> "Standortproblem":
        n, m = len(self.lager), len(self.kunden)
        if len(self.fixkosten) != n or len(self.kapazitaet) != n:
            raise ValueError(f"fixkosten/kapazitaet muessen {n} Eintraege haben")
        if len(self.bedarf) != m:
            raise ValueError(f"bedarf muss {m} Eintraege haben")
        if len(self.transport) != n or any(len(z) != m for z in self.transport):
            raise ValueError(f"transport muss {n} x {m} sein")
        if sum(self.kapazitaet) < sum(self.bedarf):
            raise ValueError(f"Gesamtkapazitaet {sum(self.kapazitaet)} deckt den "
                             f"Gesamtbedarf {sum(self.bedarf)} nicht")
        return self

    def schluessel(self, i: int, j: int) -> str:
        """Variablenname im Loesung-DTO - beide Modellbauer benutzen ihn."""
        return f"{self.lager[i]}->{self.kunden[j]}"


def beispielproblem(saat: int = 11) -> Standortproblem:
    """Sechs moegliche Lager, zwoelf Kunden - klein genug fuer beide Solver."""
    rng = np.random.default_rng(saat)
    lager = [f"Lager_{k}" for k in "ABCDEF"]
    kunden = [f"Kunde_{k:02d}" for k in range(1, 13)]
    bedarf = rng.integers(10, 60, len(kunden))
    return Standortproblem(
        lager=lager,
        kunden=kunden,
        fixkosten=rng.integers(3000, 9000, len(lager)).tolist(),
        kapazitaet=[int(bedarf.sum() * 0.45)] * len(lager),
        bedarf=bedarf.tolist(),
        transport=rng.integers(200, 1800, (len(lager), len(kunden))).tolist(),
    )


# --- 2. Der eine Baustein, der sich aendert: der Modellbauer -----------------

def baue_und_loese_mit_cpsat(problem: Standortproblem) -> Loesung:
    """CP-SAT: Bool-Variablen, ganzzahlige Koeffizienten, Minimize."""
    from ortools.sat.python import cp_model

    n, m = len(problem.lager), len(problem.kunden)
    modell = cp_model.CpModel()
    y = [modell.NewBoolVar(f"offen_{i}") for i in range(n)]
    x = {(i, j): modell.NewBoolVar(f"liefert_{i}_{j}")
         for i in range(n) for j in range(m)}

    for j in range(m):                                   # jeder Kunde genau einmal
        modell.AddExactlyOne(x[i, j] for i in range(n))
    for i in range(n):
        for j in range(m):                               # nur aus offenen Lagern
            modell.AddImplication(x[i, j], y[i])
        modell.Add(sum(problem.bedarf[j] * x[i, j] for j in range(m))
                   <= problem.kapazitaet[i] * y[i])      # Kapazitaet

    modell.Minimize(
        sum(problem.fixkosten[i] * y[i] for i in range(n))
        + sum(problem.transport[i][j] * x[i, j] for i in range(n) for j in range(m)))

    loeser = cp_model.CpSolver()
    loeser.parameters.max_time_in_seconds = ZEITLIMIT
    loeser.parameters.num_workers = 1
    loeser.parameters.random_seed = 1
    status = status_von_cpsat(loeser.Solve(modell))
    if not status.brauchbar:
        return Loesung(status=status, laufzeit=loeser.WallTime())

    return Loesung(
        status=status,
        werte={problem.schluessel(i, j): float(loeser.Value(x[i, j]))
               for i in range(n) for j in range(m)},
        zielwert=loeser.ObjectiveValue(),
        schranke=loeser.BestObjectiveBound(),
        laufzeit=loeser.WallTime())


def baue_und_loese_mit_highs(problem: Standortproblem) -> Loesung:
    """HiGHS: dieselben Restriktionen als Ungleichungszeilen einer Matrix."""
    import highspy

    n, m = len(problem.lager), len(problem.kunden)
    anzahl_x = n * m                                     # Spalten 0..n*m-1
    spalte_y = lambda i: anzahl_x + i                    # danach die y_i  # noqa: E731

    modell = highspy.Highs()
    modell.setOptionValue("output_flag", False)
    modell.setOptionValue("time_limit", ZEITLIMIT)
    modell.addVars(anzahl_x + n, np.zeros(anzahl_x + n), np.ones(anzahl_x + n))
    for spalte in range(anzahl_x + n):
        modell.changeColIntegrality(spalte, highspy.HighsVarType.kInteger)
    for i in range(n):
        modell.changeColCost(spalte_y(i), float(problem.fixkosten[i]))
        for j in range(m):
            modell.changeColCost(i * m + j, float(problem.transport[i][j]))

    for j in range(m):                                   # jeder Kunde genau einmal
        index = np.array([i * m + j for i in range(n)], dtype=np.int32)
        modell.addRow(1.0, 1.0, n, index, np.ones(n))
    for i in range(n):
        for j in range(m):                               # x_ij - y_i <= 0
            modell.addRow(-highspy.kHighsInf, 0.0, 2,
                          np.array([i * m + j, spalte_y(i)], dtype=np.int32),
                          np.array([1.0, -1.0]))
        index = np.array([i * m + j for j in range(m)] + [spalte_y(i)],
                         dtype=np.int32)                 # Kapazitaet
        werte = np.concatenate([np.array(problem.bedarf, dtype=float),
                                [-float(problem.kapazitaet[i])]])
        modell.addRow(-highspy.kHighsInf, 0.0, m + 1, index, werte)

    t0 = time.perf_counter()
    modell.run()
    laufzeit = time.perf_counter() - t0

    status = status_von_highs(modell.modelStatusToString(modell.getModelStatus()))
    if not status.brauchbar:
        return Loesung(status=status, laufzeit=laufzeit)

    loesungswerte = modell.getSolution().col_value
    info = modell.getInfo()
    return Loesung(
        status=status,
        werte={problem.schluessel(i, j): float(loesungswerte[i * m + j])
               for i in range(n) for j in range(m)},
        zielwert=info.objective_function_value,
        schranke=info.mip_dual_bound,
        laufzeit=laufzeit)


MODELLBAUER = {"cpsat": baue_und_loese_mit_cpsat, "highs": baue_und_loese_mit_highs}


# --- 3. Alles Weitere ist wieder gemeinsam -----------------------------------

def pruefe_zuordnung(problem: Standortproblem, loesung: Loesung,
                     toleranz: float = 1e-6) -> list[str]:
    """Prueft die Loesung gegen die Anforderung - ohne Solver, ohne Modell."""
    if not loesung.status.brauchbar:
        return [f"kein verwertbares Ergebnis ({loesung.status.value})"]

    n, m = len(problem.lager), len(problem.kunden)
    zuordnung = np.array([[loesung.werte[problem.schluessel(i, j)]
                           for j in range(m)] for i in range(n)])
    beanstandungen: list[str] = []

    if (np.abs(zuordnung - np.round(zuordnung)) > toleranz).any():
        beanstandungen.append("Zuordnungen sind nicht 0/1")
    zuordnung = np.round(zuordnung)

    for j, kunde in enumerate(problem.kunden):
        if abs(zuordnung[:, j].sum() - 1.0) > toleranz:
            beanstandungen.append(f"{kunde} wird {zuordnung[:, j].sum():.0f}-mal beliefert")

    beliefert = zuordnung @ np.array(problem.bedarf, dtype=float)
    for i, lagername in enumerate(problem.lager):
        if beliefert[i] > problem.kapazitaet[i] + toleranz:
            beanstandungen.append(f"{lagername}: {beliefert[i]:.0f} Paletten ueber "
                                  f"Kapazitaet {problem.kapazitaet[i]}")

    if loesung.zielwert is not None:
        offen = beliefert > toleranz
        nachgerechnet = (np.array(problem.fixkosten, dtype=float) @ offen
                         + (np.array(problem.transport, dtype=float) * zuordnung).sum())
        if abs(nachgerechnet - loesung.zielwert) > 0.5:
            beanstandungen.append(f"Zielwert {loesung.zielwert:,.0f} passt nicht zur "
                                  f"Zuordnung (nachgerechnet {nachgerechnet:,.0f})")
    return beanstandungen


def geoeffnete_lager(problem: Standortproblem, loesung: Loesung) -> list[str]:
    m = len(problem.kunden)
    return [name for i, name in enumerate(problem.lager)
            if any(loesung.werte[problem.schluessel(i, j)] > 0.5 for j in range(m))]


def loese_in_eigenem_prozess(name: str, problem: Standortproblem) -> Loesung:
    """Laesst genau einen Modellbauer in einem frischen Prozess rechnen.

    'spawn' statt des Linux-Standards 'fork': Der Kindprozess startet mit
    einem leeren Interpreter und importiert nur den Solver, den SEIN
    Modellbauer braucht. max_tasks_per_child=1 sorgt dafuer, dass der Pool
    seinen Arbeiter nicht wiederverwendet - sonst saessen beim zweiten Aufruf
    wieder beide Bibliotheken im selben Prozess.

    Hin und zurueck wandert das Domaenenmodell bzw. das Loesungs-DTO. Beide
    kennen keinen Solver, sind also serialisierbar - genau dafuer sind sie da.
    """
    with ProcessPoolExecutor(
            max_workers=1,
            mp_context=multiprocessing.get_context("spawn"),
            max_tasks_per_child=1) as pool:
        return pool.submit(MODELLBAUER[name], problem).result(timeout=300)


if __name__ == "__main__":
    problem = beispielproblem()

    # --- Beide Solver anstossen und vergleichen ---------------------------
    print("=" * 82)
    print("  DERSELBE FALL, ZWEI SOLVER - UND EIN AUSWERTUNGSCODE")
    print("=" * 82)
    print(f"Standortplanung: {len(problem.lager)} moegliche Lager, "
          f"{len(problem.kunden)} Kunden, {sum(problem.bedarf)} Paletten Bedarf.")
    print(f"Kapazitaet je Lager: {problem.kapazitaet[0]} Paletten "
          f"-> mindestens 3 Lager noetig.\n")

    loesungen: dict[str, Loesung] = {}
    for name, beschriftung in [("cpsat", "OR-Tools CP-SAT"),
                               ("highs", "HiGHS (highspy)")]:
        loesung = loesungen[name] = loese_in_eigenem_prozess(name, problem)
        beanstandungen = pruefe_zuordnung(problem, loesung)

        print(f"{beschriftung}")
        print(f"  {loesung.als_bericht()}")
        print(f"  eroeffnete Lager: {', '.join(geoeffnete_lager(problem, loesung))}")
        print(f"  Abnahmepruefung:  "
              f"{'bestanden' if not beanstandungen else beanstandungen}")

    # --- Was der Vergleich zeigt -----------------------------------------
    zielwerte = [loesung.zielwert for loesung in loesungen.values()]
    print("-" * 82)
    print(f"Zielwertdifferenz: {abs(zielwerte[0] - zielwerte[1]):.6f} EUR")

    gleich_belegt = all(
        round(loesungen["cpsat"].werte[s]) == round(loesungen["highs"].werte[s])
        for s in loesungen["cpsat"].werte)
    print(f"Identische Zuordnung: {'ja' if gleich_belegt else 'nein'}")

    assert abs(zielwerte[0] - zielwerte[1]) < 0.5, "Die Solver widersprechen sich!"
    assert all(l.status is SolverStatus.OPTIMAL for l in loesungen.values())

    print("\n" + "=" * 82)
    print("  WAS DER WECHSEL GEKOSTET HAT")
    print("=" * 82)
    print("Ausgetauscht wurde EINE Funktion. Domaenenmodell, Abnahmepruefung und")
    print("Bericht sind woertlich dieselben - sie sehen den Solver nie.")
    print()
    print("Nicht umsonst ist der Wechsel trotzdem:")
    print("  * CP-SAT rechnet ausschliesslich GANZZAHLIG. Alle Kosten sind hier")
    print("    deshalb int. Wer in Euro und Cent rechnet, skaliert vorher auf Cent -")
    print("    und muss das im Bericht wieder zuruecknehmen.")
    print("  * HiGHS braucht die Restriktionen als Matrixzeilen, CP-SAT nimmt sie")
    print("    als Ausdruecke. Das ist der Grund, warum der HiGHS-Modellbauer")
    print("    laenger ist, obwohl er dasselbe Modell beschreibt.")
    print("  * Beide Bibliotheken bringen eine eigene HiGHS-Kopie mit und lassen")
    print("    sich nicht gemeinsam importieren - daher die zwei Prozesse.")
    print()
    print("Der Ertrag: Beide beweisen denselben optimalen Zielwert, und die")
    print("Entscheidung zwischen ihnen ist eine Frage der Laufzeit geworden -")
    print("nicht eine Frage, wie viel Code man neu schreiben muss.")
    print()
    print("Verglichen wird deshalb der ZIELWERT, nicht der Plan: Gibt es mehrere")
    print("gleich teure Loesungen, darf jeder Solver eine andere davon liefern.")
    print("Hier stimmen sie zufaellig ueberein - darauf zu testen waere trotzdem")
    print("ein unzuverlaessiger Test (siehe JobShop_Intervalle.py).")
    print("=" * 82)
