#!/usr/bin/env python3

# CP_SAT_Vertretungssystem.py
"""
Kapitel CP-SAT: Dynamisches Vertretungs- und Einsatzplanungssystem mit CP-SAT.

Eigenschaften:
  * Solver-Status wird ueber SolverStatus aus or_kern.py ausgewertet - dieselbe
    Fallunterscheidung wie bei jedem anderen Solver im Buch
  * Strafkosten-Zerlegung wird ausgewiesen (Erklaerbarkeit, Kapitel Praxisfallen)
  * Fairness-Kriterium ergaenzt (Spannweite der Belastung minimieren)
  * Vorabdiagnose auf offensichtliche Unloesbarkeit
  * Abnahmepruefung des fertigen Plans OHNE den Solver zu fragen

Benoetigt: ortools, pandas, pydantic (ueber or_kern)
"""

from __future__ import annotations

from ortools.sat.python import cp_model
import pandas as pd

from or_kern import Loesung, SolverStatus, status_von_cpsat

# --- 1. Datenbasis --------------------------------------------------------
SLOTS = [1, 2, 3, 4]
FAECHER = ["Mathematik", "Physik", "Mathematik", "Informatik"]
KLASSEN = ["Klasse 8a", "Klasse 10b", "Klasse 7a", "Klasse 9c"]

PERSONAL = ["Frau_Mueller", "Herr_Schmidt", "Frau_Albrecht", "Herr_Bauer", "Frau_Koch"]

QUALIFIKATION = {
    "Frau_Mueller":  {"Mathematik", "Physik"},
    "Herr_Schmidt":  {"Informatik", "Mathematik"},
    "Frau_Albrecht": {"Physik", "Informatik"},
    "Herr_Bauer":    {"Mathematik"},
    "Frau_Koch":     {"Mathematik", "Physik", "Informatik"},
}

VORBELASTUNG = {                      # bereits geleistete Stunden heute
    "Frau_Mueller": 1, "Herr_Schmidt": 0, "Frau_Albrecht": 2,
    "Herr_Bauer": 0, "Frau_Koch": 1,
}

MAX_VERTRETUNGEN = 2
STRAFE_VORBELASTUNG = 50              # je Stunde Vorbelastung und Zuweisung
STRAFE_LOCH = 80                      # je Freistundenloch
STRAFE_UNFAIRNESS = 30                # je Einheit Spannweite der Gesamtbelastung


def pruefe_grundsaetzliche_loesbarkeit() -> bool:
    """Vorabdiagnose: Gibt es fuer jeden Slot ueberhaupt qualifiziertes Personal?"""
    ok = True
    for s, fach in enumerate(FAECHER):
        kandidaten = [p for p in PERSONAL if fach in QUALIFIKATION[p]]
        if not kandidaten:
            print(f"  UNLOESBAR: Fuer Slot {SLOTS[s]} ({fach}) gibt es niemanden.")
            ok = False
    kapazitaet = len(PERSONAL) * MAX_VERTRETUNGEN
    if kapazitaet < len(SLOTS):
        print(f"  UNLOESBAR: Kapazitaet {kapazitaet} < {len(SLOTS)} offene Stunden.")
        ok = False
    return ok


def pruefe_plan(plan: dict[int, str]) -> list[str]:
    """Prueft den fertigen Plan gegen die harten Regeln - ohne den Solver.

    Der Solver kann nur pruefen, was ihm gesagt wurde. Diese Funktion prueft
    gegen die ANFORDERUNG und benutzt dafuer bewusst keine Modellvariable
    (siehe Kapitel Praxisfallen). Leere Liste heisst bestanden.
    """
    beanstandungen: list[str] = []
    if sorted(plan) != list(range(len(SLOTS))):
        beanstandungen.append(f"nicht jede Stunde genau einmal besetzt: {sorted(plan)}")
        return beanstandungen
    for s, person in plan.items():
        if FAECHER[s] not in QUALIFIKATION[person]:
            beanstandungen.append(f"{person} ist nicht fuer {FAECHER[s]} qualifiziert")
    for person in PERSONAL:
        anzahl = sum(1 for p in plan.values() if p == person)
        if anzahl > MAX_VERTRETUNGEN:
            beanstandungen.append(f"{person} hat {anzahl} Vertretungen "
                                  f"(hoechstens {MAX_VERTRETUNGEN})")
    return beanstandungen


def plane_vertretung() -> list[dict[str, str]] | None:
    if not pruefe_grundsaetzliche_loesbarkeit():
        return None

    modell = cp_model.CpModel()

    # Entscheidungsvariablen: x[person, slot] = 1  <=>  Person uebernimmt Slot
    x = {(p, s): modell.NewBoolVar(f"zuweisung_{p}_slot{s+1}")
         for p in PERSONAL for s in range(len(SLOTS))}

    # --- HARTE NEBENBEDINGUNGEN ------------------------------------------
    # H1: Jede Stunde genau einmal besetzen
    for s in range(len(SLOTS)):
        modell.AddExactlyOne(x[p, s] for p in PERSONAL)

    # H2: Qualifikation - unqualifizierte Zuweisung ausschliessen
    for s, fach in enumerate(FAECHER):
        for p in PERSONAL:
            if fach not in QUALIFIKATION[p]:
                modell.Add(x[p, s] == 0)

    # H3: Hoechstens MAX_VERTRETUNGEN Stunden je Person
    for p in PERSONAL:
        modell.Add(sum(x[p, s] for s in range(len(SLOTS))) <= MAX_VERTRETUNGEN)

    # (H4 "nicht an zwei Orten gleichzeitig" ist durch H1 bereits erfuellt:
    #  jeder Slot hat genau eine Person, und eine Person kann pro Slot nur
    #  eine Variable auf 1 setzen.)

    # --- WEICHE ZIELE ----------------------------------------------------
    strafterme = []

    # W1: Vorbelastete Personen schonen
    strafe_last = []
    for p in PERSONAL:
        zuweisungen = sum(x[p, s] for s in range(len(SLOTS)))
        strafe_last.append(VORBELASTUNG[p] * STRAFE_VORBELASTUNG * zuweisungen)
    strafterme.extend(strafe_last)

    # W2: Freistundenloecher vermeiden (Muster: arbeitet, frei, arbeitet)
    loch_variablen = []
    for p in PERSONAL:
        for s in range(len(SLOTS) - 2):
            loch = modell.NewBoolVar(f"loch_{p}_{s}")
            # Reifizierung: loch == 1  <=>  (x_s AND NOT x_{s+1} AND x_{s+2})
            modell.AddBoolAnd([x[p, s], x[p, s + 1].Not(), x[p, s + 2]]).OnlyEnforceIf(loch)
            modell.AddBoolOr([x[p, s].Not(), x[p, s + 1], x[p, s + 2].Not()]) \
                  .OnlyEnforceIf(loch.Not())
            loch_variablen.append(loch)
            strafterme.append(loch * STRAFE_LOCH)

    # W3: Fairness - Spannweite der Gesamtbelastung minimieren
    gesamtlast = {}
    for p in PERSONAL:
        last = modell.NewIntVar(0, len(SLOTS) + max(VORBELASTUNG.values()), f"last_{p}")
        modell.Add(last == VORBELASTUNG[p] + sum(x[p, s] for s in range(len(SLOTS))))
        gesamtlast[p] = last
    max_last = modell.NewIntVar(0, 10, "max_last")
    min_last = modell.NewIntVar(0, 10, "min_last")
    modell.AddMaxEquality(max_last, list(gesamtlast.values()))
    modell.AddMinEquality(min_last, list(gesamtlast.values()))
    spannweite = modell.NewIntVar(0, 10, "spannweite")
    modell.Add(spannweite == max_last - min_last)
    strafterme.append(spannweite * STRAFE_UNFAIRNESS)

    modell.Minimize(sum(strafterme))

    # --- LOESEN -----------------------------------------------------------
    loeser = cp_model.CpSolver()
    loeser.parameters.max_time_in_seconds = 5.0
    # Reproduzierbarkeit: ein Arbeiter, fester Startwert. Ohne das kann CP-SAT
    # bei mehreren gleich guten Plaenen von Lauf zu Lauf einen anderen liefern
    # (siehe JobShop_Intervalle.py im naechsten Abschnitt).
    loeser.parameters.num_workers = 1
    loeser.parameters.random_seed = 1
    rohstatus = loeser.Solve(modell)

    # Der Rueckgabewert wird in die gemeinsame Sprache uebersetzt. Ab hier
    # sieht die Auswertung genauso aus wie bei HiGHS, GLOP oder CVXPY.
    status = status_von_cpsat(rohstatus)
    loesung = Loesung(
        status=status,
        zielwert=loeser.ObjectiveValue() if status.brauchbar else None,
        schranke=loeser.BestObjectiveBound() if status.brauchbar else None,
        laufzeit=loeser.WallTime())

    print("=" * 78)
    print("     DYNAMISCHER VERTRETUNGSPLAN (CP-SAT OPTIMIERT)")
    print("=" * 78)

    if status.modellfehler:
        print(f"Kein zulaessiger Plan moeglich (Status: {status.value}).")
        print("Die harten Regeln widersprechen sich. Naechster Schritt: eine Regel")
        print("weich machen (Kapitel Praxisfallen) oder den Konflikt eingrenzen")
        print("(Anhang Fehlerdiagnose).")
        return None
    if not status.brauchbar:
        print(f"Solver lieferte kein Ergebnis (Status: {status.value}).")
        print("Zeitlimit erhoehen oder das Modell vereinfachen.")
        return None

    guete = ("beweisbar optimal" if status is SolverStatus.OPTIMAL
             else "zulaessig, aber nicht bewiesen - Gap beachten")
    print(f"Solver-Status: {loeser.StatusName(rohstatus)} -> {status.value} ({guete})")
    print(f"{loesung.als_bericht()}\n")

    # --- Plan ausgeben ----------------------------------------------------
    zuweisung = {s: next(p for p in PERSONAL if loeser.Value(x[p, s]) == 1)
                 for s in range(len(SLOTS))}
    plan = [{
        "Stunde": f"Std. {SLOTS[s]}",
        "Klasse": KLASSEN[s],
        "Fach": FAECHER[s],
        "Vertretung": zuweisung[s],
        "Vorbelastung": f"{VORBELASTUNG[zuweisung[s]]} Std.",
    } for s in range(len(SLOTS))]
    print(pd.DataFrame(plan).to_string(index=False))

    # --- Abnahmepruefung gegen die Anforderung ----------------------------
    beanstandungen = pruefe_plan(zuweisung)
    if beanstandungen:
        print("\nABNAHMEPRUEFUNG FEHLGESCHLAGEN:")
        for beanstandung in beanstandungen:
            print(f"  - {beanstandung}")
        return None
    print("\nAbnahmepruefung: bestanden (Besetzung, Qualifikation, Hoechstzahl).")

    # --- Erklaerbarkeit: Strafkosten aufschluesseln -----------------------
    kosten_last = sum(VORBELASTUNG[p] * STRAFE_VORBELASTUNG
                      * sum(loeser.Value(x[p, s]) for s in range(len(SLOTS)))
                      for p in PERSONAL)
    kosten_loecher = sum(loeser.Value(v) for v in loch_variablen) * STRAFE_LOCH
    kosten_fairness = loeser.Value(spannweite) * STRAFE_UNFAIRNESS

    print("\n--- Woraus bestehen die Strafkosten? ---")
    print(f"  Vorbelastung geschont:  {kosten_last:5.0f} Punkte")
    print(f"  Freistundenloecher:     {kosten_loecher:5.0f} Punkte "
          f"({sum(loeser.Value(v) for v in loch_variablen)} Loch/Loecher)")
    print(f"  Fairness (Spannweite {loeser.Value(spannweite)}): {kosten_fairness:5.0f} Punkte")
    print(f"  {'Summe':<23} {loeser.ObjectiveValue():5.0f} Punkte")

    print("\n--- Auslastung nach Optimierung ---")
    for p in PERSONAL:
        heute = sum(loeser.Value(x[p, s]) for s in range(len(SLOTS)))
        print(f"  * {p:<15}: {VORBELASTUNG[p]} vorher + {heute} Vertretung "
              f"= {VORBELASTUNG[p] + heute} Stunden")

    print(f"\nSuchstatistik: {loeser.NumBranches()} Verzweigungen, "
          f"{loeser.NumConflicts()} Konflikte, {loeser.WallTime():.3f} s")
    print("=" * 78)
    return plan


if __name__ == "__main__":
    plane_vertretung()
