#!/usr/bin/env python3

# KKT_Nachweis.py
"""
Kapitel QP/NLP: Die KKT-Bedingungen numerisch nachpruefen.

Loest ein QP mit CVXPY, liest die Dualwerte aus und prueft alle vier
KKT-Bedingungen einzeln nach. Das ist zugleich eine Vorlage fuer die
Qualitaetssicherung eigener Modelle.
"""

import numpy as np
import cvxpy as cp


def baue_gueltige_kovarianz(vola, korrelationen, seed=0):
    """
    Baut Sigma = D * C * D aus Volatilitaeten und einer Korrelationsmatrix.
    Dieses Vorgehen ist konstruktionsbedingt positiv semidefinit - im
    Gegensatz zum nachtraeglichen Ueberschreiben der Diagonalen, das die
    positive Semidefinitheit zerstoeren kann.
    """
    C = np.array(korrelationen, dtype=float)
    assert np.allclose(C, C.T), "Korrelationsmatrix muss symmetrisch sein."
    assert np.allclose(np.diag(C), 1.0), "Diagonale der Korrelationsmatrix muss 1 sein."
    eigen = np.linalg.eigvalsh(C)
    assert eigen.min() > -1e-10, (
        f"Korrelationsmatrix ist nicht positiv semidefinit "
        f"(kleinster Eigenwert {eigen.min():.4f}). Solche Korrelationen sind unmoeglich.")
    D = np.diag(vola)
    return D @ C @ D


if __name__ == "__main__":
    # --- Ein kleines Portfolio-QP ----------------------------------------
    vola = np.array([0.20, 0.14, 0.30])                 # Volatilitaeten
    korr = [[1.00, 0.30, 0.10],
            [0.30, 1.00, 0.25],
            [0.10, 0.25, 1.00]]
    Sigma = baue_gueltige_kovarianz(vola, korr)
    mu = np.array([0.09, 0.05, 0.13])                   # erwartete Renditen
    lam = 3.0                                           # Risikoaversion

    eigenwerte = np.linalg.eigvalsh(Sigma)
    print("=" * 74)
    print("  KKT-BEDINGUNGEN AM PORTFOLIO-QP")
    print("=" * 74)
    print(f"Eigenwerte von Sigma: {np.round(eigenwerte, 6)}")
    print(f"  -> positiv definit: {bool(eigenwerte.min() > 0)}  "
          f"(Problem ist streng konvex, Loesung eindeutig)")
    print(f"  -> Konditionszahl:  {eigenwerte.max() / eigenwerte.min():.2f}")

    # --- Modell: min  lam/2 * w'Sigma w - mu'w   u.d.N. sum(w)=1, w>=0 ----
    n = len(mu)
    w = cp.Variable(n)
    ziel = cp.Minimize(0.5 * lam * cp.quad_form(w, Sigma) - mu @ w)
    budget = cp.sum(w) == 1
    nichtnegativ = w >= 0
    problem = cp.Problem(ziel, [budget, nichtnegativ])
    problem.solve()

    w_opt = w.value
    nu = budget.dual_value                    # Multiplikator der Gleichung
    lam_i = nichtnegativ.dual_value           # Multiplikatoren der Ungleichungen

    print(f"\nStatus: {problem.status}")
    print(f"Optimale Gewichte: {np.round(w_opt, 6)}")
    print(f"Zielwert:          {problem.value:.6f}")
    print(f"Multiplikator der Budgetgleichung (nu): {nu:.6f}")
    print(f"Multiplikatoren der w>=0-Bedingungen:   {np.round(lam_i, 6)}")

    # --- KKT-Bedingungen einzeln pruefen ---------------------------------
    print("\n--- Pruefung der vier KKT-Bedingungen ---")

    # 1. Stationaritaet:  grad f - lambda + nu*1 = 0
    #    f(w) = lam/2 w'Sigma w - mu'w   ->   grad f = lam*Sigma w - mu
    #    g_i(w) = -w_i <= 0              ->   grad g_i = -e_i
    #    h(w)   = sum(w) - 1 = 0         ->   grad h   = 1
    grad_f = lam * (Sigma @ w_opt) - mu
    stationaritaet = grad_f - lam_i + nu * np.ones(n)
    print(f"1. Stationaritaet   : max|Residuum| = {np.abs(stationaritaet).max():.2e}")

    # 2. Primale Zulaessigkeit
    print(f"2. Primal zulaessig : sum(w)-1 = {w_opt.sum()-1:.2e}, "
          f"min(w) = {w_opt.min():.2e}")

    # 3. Duale Zulaessigkeit
    print(f"3. Dual zulaessig   : min(lambda) = {lam_i.min():.2e}  (muss >= 0 sein)")

    # 4. Komplementaerer Schlupf: lambda_i * w_i = 0
    print(f"4. Kompl. Schlupf   : max|lambda_i * w_i| = "
          f"{np.abs(lam_i * w_opt).max():.2e}")

    alle_ok = (np.abs(stationaritaet).max() < 1e-6
               and abs(w_opt.sum() - 1) < 1e-8
               and w_opt.min() > -1e-8
               and lam_i.min() > -1e-8
               and np.abs(lam_i * w_opt).max() < 1e-6)
    print(f"\nAlle vier KKT-Bedingungen erfuellt: {alle_ok}")

    # --- Interpretation von nu -------------------------------------------
    print("\n--- Was bedeutet nu? ---")
    print("nu ist der Schattenpreis des Budgets: Um so viel aendert sich der")
    print("Zielwert, wenn man statt 100 % nur 99 % investieren duerfte.")
    problem2 = cp.Problem(cp.Minimize(0.5 * lam * cp.quad_form(w, Sigma) - mu @ w),
                          [cp.sum(w) == 1.01, w >= 0])
    problem2.solve()
    print(f"  Vorhergesagt (nu * 0.01): {nu * 0.01:+.6f}")
    print(f"  Tatsaechlich gemessen:    {problem2.value - problem.value:+.6f}")
    print("=" * 74)
