#!/usr/bin/env python3

# QuantitativeTradingEngine.py
"""
Kapitel Handelsmaschine: Vollstaendige quantitative Handelsmaschine mit
Walk-Forward-Backtest und OR-Optimierungsschicht.

Eigenschaften:
  * Rebalancing-Termine sind garantiert Handelstage; die Zahl der
    tatsaechlich ausgefuehrten Umschichtungen wird ausgewiesen
  * Spaltenreihenfolge erzwungen
  * Einheiten konsistent (taegliche Groessen im Modell)
  * Szenario-Nebenbedingungen vektorisiert
  * Lookahead-Selbsttest eingebaut
  * Zusaetzliche Kennzahlen: Calmar, Trefferquote, Turnover, Kostenanteil
"""

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)

HANDELSTAGE = 252


class QuantitativeHandelsmaschine:
    """Walk-Forward-Backtest mit CVaR-Optimierung und Transaktionskosten."""

    def __init__(self, universum, benchmark="SPY", jahre=5,
                 startkapital=100_000.0, gebuehrensatz=0.0015,
                 max_gewicht=0.20, alpha=0.95, risikoaversion=1.2,
                 warmup=HANDELSTAGE):
        self.universum = list(universum)
        self.benchmark = benchmark
        self.jahre = jahre
        self.startkapital = startkapital
        self.gebuehrensatz = gebuehrensatz
        self.max_gewicht = max_gewicht
        self.alpha = alpha
        self.risikoaversion = risikoaversion
        self.warmup = warmup

        self.kurse = None
        self.benchmark_kurse = None
        self.renditen = None
        self.rebalancing_termine = None

    # -- 1. Datenpipeline ------------------------------------------------
    def lade_daten(self):
        ende = pd.Timestamp.today().normalize()
        start = ende - pd.DateOffset(years=self.jahre)
        alle = self.universum + [self.benchmark]

        print(f"Lade {len(self.universum)} Titel + Benchmark ({self.benchmark}) "
              f"ueber {self.jahre} Jahre ...")
        roh = yf.download(alle, start=start, end=ende, auto_adjust=True, progress=False)
        if roh.empty:
            raise SystemExit("Download fehlgeschlagen (Netz/Ticker/Rate-Limit pruefen).")

        if isinstance(roh.columns, pd.MultiIndex):
            daten = roh["Close"][alle].dropna()      # erzwingt eigene Spaltenreihenfolge
        else:
            daten = roh[["Close"]].dropna()
            daten.columns = alle
        assert list(daten.columns) == alle, "Spaltenreihenfolge weicht ab!"

        self.kurse = daten[self.universum]
        self.benchmark_kurse = daten[self.benchmark]
        self.renditen = self.kurse.pct_change().dropna()

        # --- erster HANDELStag je Monat, nicht Kalendermonatsanfang -----
        monat = self.kurse.index.to_period("M")
        self.rebalancing_termine = self.kurse.index[~monat.duplicated()]
        assert self.rebalancing_termine.isin(self.kurse.index).all(), \
            "Rebalancing-Termine sind keine Handelstage!"

        print(f"  {len(self.kurse)} Handelstage, "
              f"{len(self.rebalancing_termine)} Rebalancing-Termine "
              f"(alle sind Handelstage)")

    # -- 2. Signal --------------------------------------------------------
    def alpha_signal(self, stichtag, rueckblick=HANDELSTAGE):
        """
        12-1-Momentum: Kursentwicklung der letzten 12 Monate unter Auslassung
        des letzten Monats (kurzfristige Umkehr herausrechnen).
        Rueckgabe: TAEGLICHE Renditeerwartung.
        """
        historie = self.kurse.loc[:stichtag]
        if len(historie) < rueckblick:
            return np.zeros(len(self.universum))

        p_alt = historie.iloc[-rueckblick]
        p_neu = historie.iloc[-21]                    # vor ca. einem Monat
        momentum = (p_neu / p_alt - 1.0).values

        # Z-Score, damit das Signal skalenfrei wird
        streuung = momentum.std()
        score = (momentum - momentum.mean()) / (streuung + 1e-9)

        # In eine taegliche Renditeerwartung uebersetzen:
        # Basis 4 % p.a. plus 6 % p.a. je Standardabweichung Momentum.
        mu_jaehrlich = np.maximum(0.04 + 0.06 * score, 0.0)
        return mu_jaehrlich / HANDELSTAGE

    # -- 3./4. Risikomodell und Optimierung -------------------------------
    def optimiere(self, stichtag, w_alt, rueckblick=HANDELSTAGE):
        historie = self.renditen.loc[:stichtag].iloc[-rueckblick:]
        if len(historie) < 100:
            return w_alt, False

        R = historie.values
        S, N = R.shape
        mu_taeglich = self.alpha_signal(stichtag, rueckblick)

        w = cp.Variable(N, nonneg=True)
        gamma = cp.Variable()
        u = cp.Variable(S, nonneg=True)

        cvar = gamma + (1.0 / (S * (1.0 - self.alpha))) * cp.sum(u)
        kosten = self.gebuehrensatz * cp.norm1(w - w_alt)
        ziel = cp.Maximize(mu_taeglich @ w - self.risikoaversion * cvar - kosten)

        bedingungen = [cp.sum(w) == 1,
                       w <= self.max_gewicht,
                       u >= -(R @ w) - gamma]          # vektorisiert
        problem = cp.Problem(ziel, bedingungen)
        try:
            problem.solve(solver=cp.CLARABEL)
        except cp.error.SolverError:
            return w_alt, False

        if problem.status not in ("optimal", "optimal_inaccurate") or w.value is None:
            return w_alt, False

        gewichte = np.maximum(w.value, 0.0)
        return gewichte / gewichte.sum(), True

    # -- 5. Backtest -------------------------------------------------------
    def backtest(self):
        print("Starte Walk-Forward-Backtest ...")
        N = len(self.universum)
        gewichte = np.ones(N) / N

        werte = [self.startkapital]
        datumsliste = [self.renditen.index[self.warmup]]
        gewichtshistorie = []
        geplant = versucht = ausgefuehrt = 0
        kosten_gesamt = 0.0
        turnover_gesamt = 0.0

        for t in range(self.warmup, len(self.renditen) - 1):
            heute = self.renditen.index[t]
            morgen = self.renditen.index[t + 1]

            ist_termin = heute in self.rebalancing_termine or t == self.warmup
            if ist_termin:
                geplant += 1
                versucht += 1
                neue_gewichte, erfolg = self.optimiere(heute, gewichte)
                if erfolg:
                    ausgefuehrt += 1
                    turnover = float(np.abs(neue_gewichte - gewichte).sum())
                    gebuehr = werte[-1] * turnover * self.gebuehrensatz
                    werte[-1] -= gebuehr
                    kosten_gesamt += gebuehr
                    turnover_gesamt += turnover
                    gewichte = neue_gewichte
                    gewichtshistorie.append((heute, gewichte.copy()))

            # Rendite von MORGEN auf die HEUTE festgelegten Gewichte anwenden
            tagesrenditen = self.renditen.iloc[t + 1].values
            portfoliorendite = float(gewichte @ tagesrenditen)
            werte.append(werte[-1] * (1.0 + portfoliorendite))
            datumsliste.append(morgen)

            # Gewichte driften mit den Kursen weiter
            gedriftet = gewichte * (1.0 + tagesrenditen)
            gewichte = gedriftet / gedriftet.sum()

        print(f"  Rebalancing: {geplant} geplant, {ausgefuehrt} ausgefuehrt "
              f"({ausgefuehrt/max(geplant,1)*100:.0f} %)")
        if ausgefuehrt < geplant:
            print(f"  WARNUNG: {geplant - ausgefuehrt} Optimierungen sind "
                  f"fehlgeschlagen - Ursache pruefen!")
        print(f"  Kumulierter Turnover: {turnover_gesamt*100:.0f} % | "
              f"Transaktionskosten insgesamt: {kosten_gesamt:,.2f} EUR")

        ergebnis = pd.DataFrame({"Portfolio": werte}, index=datumsliste)
        benchmark = self.benchmark_kurse.loc[datumsliste]
        ergebnis["Benchmark"] = benchmark / benchmark.iloc[0] * self.startkapital
        return ergebnis, gewichtshistorie, kosten_gesamt

    # -- 6. Auswertung -----------------------------------------------------
    @staticmethod
    def kennzahlen(reihe, risikofrei=0.02):
        renditen = reihe.pct_change().dropna()
        jahre = (reihe.index[-1] - reihe.index[0]).days / 365.25
        cagr = (reihe.iloc[-1] / reihe.iloc[0]) ** (1 / jahre) - 1
        vola = renditen.std() * np.sqrt(HANDELSTAGE)
        sharpe = (cagr - risikofrei) / vola if vola > 0 else np.nan
        drawdown = ((reihe - reihe.cummax()) / reihe.cummax()).min()
        calmar = cagr / abs(drawdown) if drawdown < 0 else np.nan
        trefferquote = float((renditen > 0).mean())
        return {"CAGR": cagr, "Volatilitaet": vola, "Sharpe": sharpe,
                "Max Drawdown": drawdown, "Calmar": calmar,
                "Trefferquote": trefferquote, "Jahre": jahre}

    def bericht(self, ergebnis, kosten_gesamt):
        k_port = self.kennzahlen(ergebnis["Portfolio"])
        k_bench = self.kennzahlen(ergebnis["Benchmark"])

        print("\n" + "=" * 84)
        print("          LEISTUNGSREPORT DER QUANTITATIVEN ENGINE")
        print("=" * 84)
        print(f"Zeitraum: {ergebnis.index[0]:%Y-%m-%d} bis {ergebnis.index[-1]:%Y-%m-%d} "
              f"({k_port['Jahre']:.1f} Jahre) | Startkapital "
              f"{self.startkapital:,.0f} EUR\n")

        zeilen = [
            ("Endkapital", f"{ergebnis['Portfolio'].iloc[-1]:,.0f} EUR",
             f"{ergebnis['Benchmark'].iloc[-1]:,.0f} EUR"),
            ("CAGR (Jahresrendite)", f"{k_port['CAGR']*100:6.2f} %",
             f"{k_bench['CAGR']*100:6.2f} %"),
            ("Volatilitaet p.a.", f"{k_port['Volatilitaet']*100:6.2f} %",
             f"{k_bench['Volatilitaet']*100:6.2f} %"),
            ("Sharpe Ratio (rf=2 %)", f"{k_port['Sharpe']:6.2f}",
             f"{k_bench['Sharpe']:6.2f}"),
            ("Maximum Drawdown", f"{k_port['Max Drawdown']*100:6.2f} %",
             f"{k_bench['Max Drawdown']*100:6.2f} %"),
            ("Calmar Ratio", f"{k_port['Calmar']:6.2f}", f"{k_bench['Calmar']:6.2f}"),
            ("Trefferquote (Tage)", f"{k_port['Trefferquote']*100:6.1f} %",
             f"{k_bench['Trefferquote']*100:6.1f} %"),
        ]
        print(pd.DataFrame(zeilen, columns=["Kennzahl", "OR-Strategie", "Benchmark"])
              .to_string(index=False))

        anteil = kosten_gesamt / self.startkapital
        print(f"\nTransaktionskosten: {kosten_gesamt:,.0f} EUR "
              f"= {anteil*100:.2f} % des Startkapitals "
              f"= {anteil/k_port['Jahre']*100:.2f} % p.a.")
        print("Ohne diese Kosten waere die ausgewiesene Rendite entsprechend hoeher -")
        print("genau deshalb gehoeren sie in den Backtest und nicht daneben.")
        print("=" * 84)
        return k_port, k_bench

    def zeichne(self, ergebnis):
        plt.figure(figsize=(11, 6))
        plt.plot(ergebnis.index, ergebnis["Portfolio"], "b-", lw=2.2,
                 label="OR-Handelsmaschine")
        plt.plot(ergebnis.index, ergebnis["Benchmark"], "k--", lw=1.5, alpha=0.75,
                 label=f"Benchmark ({self.benchmark})")
        plt.title("Walk-Forward-Backtest: OR-Strategie gegen Benchmark", fontsize=12)
        plt.xlabel("Datum")
        plt.ylabel("Depotwert in EUR")
        plt.grid(True, linestyle=":", alpha=0.6)
        plt.legend(loc="upper left")
        plt.tight_layout()
        ziel = os.path.join(OUTPUT_DIR, "trading_engine_backtest.png")
        plt.savefig(ziel, dpi=150)
        print(f"Chart gespeichert unter '{ziel}'")


def lookahead_selbsttest(maschine):
    """
    Prueft, dass das Signal zum Zeitpunkt t NUR Daten bis t verwendet.
    Methode: kuenstlich die Zukunft veraendern und pruefen, ob sich das
    Signal von heute dadurch aendert. Tut es das, liegt Lookahead vor.
    """
    stichtag = maschine.kurse.index[len(maschine.kurse) // 2]
    signal_original = maschine.alpha_signal(stichtag)

    original_kurse = maschine.kurse.copy()
    zukunft = maschine.kurse.index > stichtag
    maschine.kurse.loc[zukunft] *= 3.0            # Zukunft massiv veraendern
    signal_manipuliert = maschine.alpha_signal(stichtag)
    maschine.kurse = original_kurse               # zuruecksetzen

    identisch = np.allclose(signal_original, signal_manipuliert)
    print(f"Lookahead-Selbsttest: Signal bleibt bei manipulierter Zukunft "
          f"{'UNVERAENDERT (gut)' if identisch else 'VERAENDERT - LOOKAHEAD-BIAS!'}")
    assert identisch, "Das Signal verwendet Zukunftsdaten!"


if __name__ == "__main__":
    UNIVERSUM = ["AAPL", "MSFT", "NVDA", "GOOGL",
                 "JNJ", "UNH", "PFE",
                 "JPM", "BAC", "GS",
                 "XOM", "CVX",
                 "PG", "KO", "COST"]

    maschine = QuantitativeHandelsmaschine(UNIVERSUM, benchmark="SPY", jahre=5)
    maschine.lade_daten()
    lookahead_selbsttest(maschine)
    ergebnis, gewichte, kosten = maschine.backtest()
    maschine.bericht(ergebnis, kosten)
    maschine.zeichne(ergebnis)
