#!/usr/bin/env python3

# Markowitz_CVXPY.py
"""
Kapitel Markowitz: Markowitz-Mean-Variance-Optimierung mit CVXPY.
Enthaelt GMV, Maximum Sharpe (Korn-Transformation), Effizienzgrenze
und institutionelle Restriktionen.

Eigenschaften:
  * Spaltenreihenfolge erzwungen; Sektoren ueber NAMEN statt Indizes
  * Effizienzgrenze bis zur TATSAECHLICH erreichbaren Maximalrendite
    (nicht etwa max(mu)*0.95 - das waere unter Restriktionen oft
     unerreichbar und die Kurve wuerde stillschweigend vorzeitig abbrechen)
  * Pruefung aller Restriktionen nach dem Loesen
  * Vergleich mit Gleichgewichtung als Realitaetscheck
"""

import os

import cvxpy as cp
import numpy as np
import pandas as pd
import yfinance as yf
from sklearn.covariance import LedoitWolf
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt

OUTPUT_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "output")
os.makedirs(OUTPUT_DIR, exist_ok=True)

TICKER = ["AAPL", "MSFT", "NVDA", "AMZN", "JNJ", "PFE", "JPM", "GS", "XOM", "CVX"]

# Sektoren ueber Namen definiert, nicht ueber Positionen
SEKTOREN = {
    "Technologie": ["AAPL", "MSFT", "NVDA", "AMZN"],
    "Gesundheit":  ["JNJ", "PFE"],
    "Finanzen":    ["JPM", "GS"],
    "Energie":     ["XOM", "CVX"],
}
SEKTORGRENZEN = {"Technologie": 0.35, "Gesundheit": 0.40,
                 "Finanzen": 0.40, "Energie": 0.40}

MAX_GEWICHT = 0.20      # hoechstens 20 % je Einzeltitel
RISIKOFREI = 0.03       # 3 % p.a.
HANDELSTAGE = 252


def lade_daten():
    """Laedt Kurse und garantiert die Spaltenreihenfolge."""
    ende = pd.Timestamp.today().normalize()
    start = ende - pd.DateOffset(years=2)
    roh = yf.download(TICKER, start=start, end=ende, auto_adjust=True, progress=False)
    if roh.empty:
        raise SystemExit("Download fehlgeschlagen (Netz, Ticker oder Rate-Limit pruefen).")

    if isinstance(roh.columns, pd.MultiIndex):
        kurse = roh["Close"][TICKER].dropna()      # [TICKER] erzwingt die Reihenfolge
    else:
        kurse = roh[["Close"]].dropna()
        kurse.columns = TICKER

    assert list(kurse.columns) == TICKER, "Spaltenreihenfolge weicht ab!"
    renditen = kurse.pct_change().dropna()
    mu = renditen.mean().values * HANDELSTAGE
    sigma = LedoitWolf().fit(renditen.values).covariance_ * HANDELSTAGE
    return kurse, renditen, mu, sigma


def sektor_indizes():
    """Uebersetzt Sektornamen einmalig in Positionsindizes - mit Pruefung."""
    ergebnis = {}
    for sektor, titel in SEKTOREN.items():
        fehlend = [t for t in titel if t not in TICKER]
        if fehlend:
            raise ValueError(f"Sektor '{sektor}': Ticker {fehlend} nicht im Universum.")
        ergebnis[sektor] = [TICKER.index(t) for t in titel]
    return ergebnis


def basis_restriktionen(w, skala=None):
    """
    Standardrestriktionen. Bei der Korn-Transformation muessen ALLE Grenzen
    mit kappa mitskaliert werden - dafuer dient der Parameter 'skala'.
    """
    eins = 1.0 if skala is None else skala
    idx = sektor_indizes()
    bedingungen = [w >= 0, w <= MAX_GEWICHT * eins]
    for sektor, positionen in idx.items():
        bedingungen.append(cp.sum(w[positionen]) <= SEKTORGRENZEN[sektor] * eins)
    return bedingungen


def loese_gmv(sigma):
    """Global Minimum Variance: minimiere Risiko, ignoriere Rendite."""
    w = cp.Variable(len(sigma))
    problem = cp.Problem(cp.Minimize(0.5 * cp.quad_form(w, sigma)),
                         [cp.sum(w) == 1] + basis_restriktionen(w))
    problem.solve()
    if problem.status not in ("optimal", "optimal_inaccurate"):
        raise SystemExit(f"GMV nicht loesbar: {problem.status}")
    return w.value


def loese_max_sharpe(mu, sigma, r_f):
    """Maximum Sharpe Ratio ueber die Korn-Transformation."""
    n = len(mu)
    ueberrendite = mu - r_f
    if np.all(ueberrendite <= 0):
        raise SystemExit("Kein Titel schlaegt den risikofreien Zins - "
                         "Max-Sharpe-Portfolio existiert nicht.")

    y = cp.Variable(n)
    kappa = cp.Variable(nonneg=True)
    bedingungen = ([ueberrendite @ y == 1, cp.sum(y) == kappa]
                   + basis_restriktionen(y, skala=kappa))
    problem = cp.Problem(cp.Minimize(0.5 * cp.quad_form(y, sigma)), bedingungen)
    problem.solve()
    if problem.status not in ("optimal", "optimal_inaccurate"):
        raise SystemExit(f"Max-Sharpe nicht loesbar: {problem.status}")
    return y.value / kappa.value


def max_erreichbare_rendite(mu, sigma):
    """
    Achtung: max(mu)*0.95 als obere Grenze der Frontier anzunehmen, ist unter
    Positions- und Sektorgrenzen oft unerreichbar; die Kurve wuerde dann
    stillschweigend abbrechen. Hier wird die tatsaechliche Obergrenze berechnet.
    """
    w = cp.Variable(len(mu))
    problem = cp.Problem(cp.Maximize(mu @ w), [cp.sum(w) == 1] + basis_restriktionen(w))
    problem.solve()
    return float(problem.value)


def berechne_frontier(mu, sigma, ret_min, ret_max, punkte=40):
    """Effizienzgrenze durch Variation der Mindestrendite."""
    n = len(mu)
    w = cp.Variable(n)
    ziel_rendite = cp.Parameter()
    problem = cp.Problem(
        cp.Minimize(0.5 * cp.quad_form(w, sigma)),
        [cp.sum(w) == 1, mu @ w >= ziel_rendite] + basis_restriktionen(w))

    volas, renditen, uebersprungen = [], [], 0
    for ziel in np.linspace(ret_min, ret_max, punkte):
        ziel_rendite.value = ziel
        problem.solve()
        if problem.status in ("optimal", "optimal_inaccurate"):
            volas.append(float(np.sqrt(w.value @ sigma @ w.value)))
            renditen.append(float(mu @ w.value))
        else:
            uebersprungen += 1
    if uebersprungen:
        print(f"  Hinweis: {uebersprungen} Zielrenditen waren nicht erreichbar.")
    return np.array(volas), np.array(renditen)


def pruefe_restriktionen(w, bezeichnung):
    """Nach dem Loesen: haelt die Loesung wirklich alle Regeln ein?"""
    assert abs(w.sum() - 1) < 1e-6, f"{bezeichnung}: Summe != 1"
    assert w.min() > -1e-6, f"{bezeichnung}: negatives Gewicht"
    assert w.max() < MAX_GEWICHT + 1e-6, f"{bezeichnung}: Positionsgrenze verletzt"
    for sektor, positionen in sektor_indizes().items():
        anteil = w[positionen].sum()
        assert anteil < SEKTORGRENZEN[sektor] + 1e-6, \
            f"{bezeichnung}: Sektor {sektor} bei {anteil:.3f} ueber Grenze"


def kennzahlen(w, mu, sigma, r_f):
    rendite = float(mu @ w)
    vola = float(np.sqrt(w @ sigma @ w))
    return rendite, vola, (rendite - r_f) / vola


if __name__ == "__main__":
    kurse, renditen, mu, sigma = lade_daten()
    n = len(TICKER)

    w_gmv = loese_gmv(sigma)
    w_sharpe = loese_max_sharpe(mu, sigma, RISIKOFREI)
    w_gleich = np.ones(n) / n

    pruefe_restriktionen(w_gmv, "GMV")
    pruefe_restriktionen(w_sharpe, "Max Sharpe")

    print("=" * 88)
    print("        ERGEBNISSE DER MEAN-VARIANCE-OPTIMIERUNG")
    print("=" * 88)
    print(f"Datenbasis: {len(renditen)} Handelstage, {n} Titel, "
          f"risikofreier Zins {RISIKOFREI*100:.1f} %")
    print(f"Restriktionen: max. {MAX_GEWICHT*100:.0f} % je Titel, "
          f"Sektorgrenzen {SEKTORGRENZEN}\n")

    print(f"{'Portfolio':<26} {'Rendite':>10} {'Volatilitaet':>13} {'Sharpe':>9}")
    print("-" * 88)
    for name, w in [("Global Minimum Variance", w_gmv),
                    ("Maximum Sharpe Ratio", w_sharpe),
                    ("Gleichgewichtung (1/N)", w_gleich)]:
        r, v, sr = kennzahlen(w, mu, sigma, RISIKOFREI)
        print(f"{name:<26} {r*100:>9.2f} % {v*100:>12.2f} % {sr:>9.2f}")

    # --- Gewichte je Titel ------------------------------------------------
    print("\n--- Optimierte Portfoliogewichte ---")
    sektor_je_titel = {t: s for s, titel in SEKTOREN.items() for t in titel}
    tabelle = pd.DataFrame({
        "Ticker": TICKER,
        "Sektor": [sektor_je_titel[t] for t in TICKER],
        "Rendite p.a.": [f"{r*100:+6.1f} %" for r in mu],
        "Vola p.a.": [f"{np.sqrt(sigma[i, i])*100:5.1f} %" for i in range(n)],
        "GMV": [f"{w*100:5.1f} %" for w in w_gmv],
        "Max Sharpe": [f"{w*100:5.1f} %" for w in w_sharpe],
    })
    print(tabelle.to_string(index=False))

    print("\n--- Sektoraufteilung (Kontrolle) ---")
    print(f"{'Sektor':<14} {'Grenze':>8} {'GMV':>9} {'Max Sharpe':>12}")
    for sektor, positionen in sektor_indizes().items():
        print(f"{sektor:<14} {SEKTORGRENZEN[sektor]*100:>7.0f} % "
              f"{w_gmv[positionen].sum()*100:>8.1f} % "
              f"{w_sharpe[positionen].sum()*100:>11.1f} %")

    # --- Effizienzgrenze --------------------------------------------------
    ret_gmv = float(mu @ w_gmv)
    ret_max = max_erreichbare_rendite(mu, sigma)
    print(f"\nEffizienzgrenze von {ret_gmv*100:.2f} % bis {ret_max*100:.2f} % "
          f"(unter Restriktionen tatsaechlich erreichbar)")
    print(f"  Zum Vergleich: bester Einzeltitel {mu.max()*100:.2f} % - "
          f"durch die Grenzen nicht erreichbar.")
    volas, rets = berechne_frontier(mu, sigma, ret_gmv, ret_max)

    # --- Diagramm ---------------------------------------------------------
    plt.figure(figsize=(11, 6.5))
    plt.plot(volas * 100, rets * 100, "b-", lw=2.5, label="Effizienzgrenze (restringiert)")

    for w, farbe, marker, groesse, name in [
            (w_gmv, "green", "o", 150, "GMV"),
            (w_sharpe, "red", "*", 260, "Max Sharpe"),
            (w_gleich, "purple", "D", 110, "Gleichgewichtung")]:
        r, v, sr = kennzahlen(w, mu, sigma, RISIKOFREI)
        plt.scatter([v * 100], [r * 100], color=farbe, marker=marker, s=groesse,
                    zorder=5, label=f"{name} (SR={sr:.2f})")

    # Kapitalmarktlinie durch den risikofreien Zins und das Tangentialportfolio
    r_s, v_s, _ = kennzahlen(w_sharpe, mu, sigma, RISIKOFREI)
    x_linie = np.array([0, v_s * 1.25])
    plt.plot(x_linie * 100, (RISIKOFREI + (r_s - RISIKOFREI) / v_s * x_linie) * 100,
             color="orange", ls=":", lw=2, label="Kapitalmarktlinie")

    for i, t in enumerate(TICKER):
        plt.scatter(np.sqrt(sigma[i, i]) * 100, mu[i] * 100, color="gray", alpha=0.5, s=40)
        plt.annotate(t, (np.sqrt(sigma[i, i]) * 100 + 0.4, mu[i] * 100), fontsize=8)

    plt.title("Markowitz-Effizienzgrenze mit Sektor- und Positionsgrenzen", fontsize=12)
    plt.xlabel("Annualisierte Volatilitaet [%]")
    plt.ylabel("Annualisierte erwartete Rendite [%]")
    plt.grid(True, linestyle=":", alpha=0.6)
    plt.legend(loc="best")
    plt.tight_layout()
    ziel = os.path.join(OUTPUT_DIR, "markowitz_efficient_frontier.png")
    plt.savefig(ziel, dpi=150)
    print(f"\nDiagramm gespeichert unter '{ziel}'")
    print("=" * 88)
