"""
Agente HRP de Leonex - Hierarchical Risk Parity.

Implementacion del algoritmo de Marcos Lopez de Prado (Advances in Financial
Machine Learning, Wiley 2018, capitulo 16). Asigna pesos en una cartera SIN
tener que invertir la matriz de covarianzas (que es donde Markowitz se rompe
con activos correlacionados).

REFERENCIAS BIBLIOGRAFICAS
- Lopez de Prado, M. (2016). "Building Diversified Portfolios that Outperform
  Out of Sample". Journal of Portfolio Management 42(4). Paper original del HRP.
- Lopez de Prado, M. (2018). Advances in Financial Machine Learning. Wiley.
  Capitulo 16 (HRP) con pseudocodigo completo y comparaciones vs Markowitz.
- Dalio, R. (2017). Principles. Bridgewater All Weather. Inspira el espiritu
  risk-parity (igualar contribuciones al riesgo, no al capital), aunque HRP
  es la version "matematicamente honesta" del concepto.
- Markowitz, H. (1952). "Portfolio Selection". Journal of Finance. La pieza
  que HRP REEMPLAZA por su inestabilidad ante correlaciones altas.

Tres pasos:

    1. Tree clustering por distancia de correlacion:
       d(i,j) = sqrt(0.5 * (1 - rho(i,j)))
       Genera un dendrograma usando linkage single-link.

    2. Quasi-diagonalizacion:
       Reordena los activos para que los similares queden juntos en la
       matriz de covarianzas. Esto agrupa los riesgos correlacionados.

    3. Recursive bisection:
       Divide recursivamente el universo en dos mitades por longitud.
       Asigna pesos a cada mitad proporcionales a la inversa de su
       varianza interna (los clusters menos volatiles reciben mas peso).

Ventajas vs Markowitz:
    - No requiere invertir matrices (resistente a errores de estimacion).
    - Maneja activos altamente correlacionados sin explotar.
    - Produce carteras mas estables out-of-sample.

Uso:
    python agents/agente_hrp.py                          # tickers del active_universe
    python agents/agente_hrp.py --tickers SPY,QQQ,AAPL   # tickers especificos
    python agents/agente_hrp.py --window 120             # ventana retornos
    python agents/agente_hrp.py --plan                   # tickers del plan del dia
"""

from __future__ import annotations

import argparse
import json
import logging
import sqlite3
import sys
from dataclasses import asdict, dataclass, field
from datetime import datetime
try:
    from datetime import UTC  # Python 3.11+
except ImportError:  # Python 3.10
    from datetime import timezone
    UTC = timezone.utc
from pathlib import Path
from typing import Optional

import numpy as np
import pandas as pd


PROJECT_ROOT = Path(__file__).resolve().parents[1]
DATA_DIR = PROJECT_ROOT / "data"
LOGS_DIR = PROJECT_ROOT / "logs"
DASHBOARD_DATA_DIR = PROJECT_ROOT / "dashboard" / "data"
DB_PATH = DATA_DIR / "Leonex.sqlite"
REPORT_PATH = DASHBOARD_DATA_DIR / "hrp_report.json"

DEFAULT_WINDOW = 60                # dias de retornos para covarianza
MIN_TICKERS = 2                    # con 1 solo no tiene sentido HRP


@dataclass
class HrpWeight:
    ticker: str
    weight: float
    rank: int                       # 1 = mayor peso
    annual_vol: float               # volatilidad anualizada del activo


@dataclass
class HrpReport:
    generated_at: str
    window_days: int
    n_tickers: int
    method: str                     # "HRP" | "single_asset" | "fallback_equal_weight"
    weights: list[HrpWeight] = field(default_factory=list)
    tickers_requested: list[str] = field(default_factory=list)
    tickers_used: list[str] = field(default_factory=list)
    skipped: dict = field(default_factory=dict)   # ticker → razon
    correlation_avg: Optional[float] = None       # correlacion media (info)
    note: str = ""


# ── Algoritmo HRP ─────────────────────────────────────────────────────────

def _correlation_distance(corr: np.ndarray) -> np.ndarray:
    """Convierte matriz de correlacion en matriz de distancia."""
    # d(i,j) = sqrt(0.5 * (1 - rho)) ∈ [0, 1]
    d = np.sqrt(np.clip(0.5 * (1 - corr), 0, 2))
    # Diagonal a 0 para evitar self-distance > 0 por ruido numerico
    np.fill_diagonal(d, 0.0)
    return d


def _get_quasi_diag(link: np.ndarray) -> list[int]:
    """Lopez de Prado AFML §16.4: recupera el orden de hojas del dendrograma."""
    link = link.astype(int)
    sort_ix = pd.Series([link[-1, 0], link[-1, 1]])
    num_items = int(link[-1, 3])
    while sort_ix.max() >= num_items:
        sort_ix.index = range(0, sort_ix.shape[0] * 2, 2)
        df0 = sort_ix[sort_ix >= num_items]
        i = df0.index
        j = df0.values - num_items
        sort_ix[i] = link[j, 0]
        df1 = pd.Series(link[j, 1], index=i + 1)
        sort_ix = pd.concat([sort_ix, df1])
        sort_ix = sort_ix.sort_index()
        sort_ix.index = range(sort_ix.shape[0])
    return sort_ix.tolist()


def _cluster_variance(cov: pd.DataFrame, tickers: list[str]) -> float:
    """Varianza de un cluster usando pesos por varianza inversa (IVP)."""
    c = cov.loc[tickers, tickers].values
    diag = np.diag(c)
    if (diag <= 0).any():
        return 1e9
    ivp = 1.0 / diag
    ivp = ivp / ivp.sum()
    return float(ivp @ c @ ivp)


def _recursive_bisection(cov: pd.DataFrame, sorted_tickers: list[str]) -> pd.Series:
    """Recursive bisection (AFML §16.5)."""
    weights = pd.Series(1.0, index=sorted_tickers, dtype=float)
    clusters = [sorted_tickers]
    while clusters:
        new_clusters: list[list[str]] = []
        for cluster in clusters:
            if len(cluster) <= 1:
                continue
            mid = len(cluster) // 2
            left = cluster[:mid]
            right = cluster[mid:]
            var_left = _cluster_variance(cov, left)
            var_right = _cluster_variance(cov, right)
            if var_left + var_right == 0:
                alpha = 0.5
            else:
                alpha = 1 - var_left / (var_left + var_right)
            weights[left] *= alpha
            weights[right] *= (1 - alpha)
            new_clusters.append(left)
            new_clusters.append(right)
        clusters = new_clusters
    return weights


def hrp_weights(returns: pd.DataFrame) -> tuple[pd.Series, float]:
    """
    Calcula pesos HRP a partir de un DataFrame de retornos diarios.
    Devuelve (serie weights con suma 1, correlacion media).
    """
    if returns.shape[1] < 2:
        # Caso degenerado: 1 solo ticker
        w = pd.Series([1.0], index=returns.columns)
        return w, 1.0

    # 1) Limpiamos NaN
    returns = returns.dropna(axis=1, thresh=int(len(returns) * 0.6))
    if returns.shape[1] < 2:
        w = pd.Series([1.0], index=returns.columns)
        return w, 1.0
    returns = returns.fillna(method="ffill").fillna(method="bfill")

    # 2) Covarianza y correlacion
    cov = returns.cov()
    corr = returns.corr().clip(-1, 1)

    # 3) Tree clustering: dendrograma con single linkage
    from scipy.cluster.hierarchy import linkage
    from scipy.spatial.distance import squareform

    dist = _correlation_distance(corr.values)
    # squareform necesita una matriz simetrica con 0 en diagonal
    np.fill_diagonal(dist, 0.0)
    # Garantizamos simetria estricta para squareform
    dist = (dist + dist.T) / 2.0
    cond = squareform(dist, checks=False)
    link = linkage(cond, method="single")

    # 4) Quasi-diagonalizacion
    sorted_idx = _get_quasi_diag(link)
    sorted_tickers = [returns.columns[i] for i in sorted_idx]

    # 5) Recursive bisection
    weights = _recursive_bisection(cov, sorted_tickers)
    weights = weights / weights.sum()  # normalizar

    # Correlacion media (estadistica info)
    upper = corr.values[np.triu_indices(len(corr), k=1)]
    corr_avg = float(np.mean(upper)) if len(upper) > 0 else 1.0

    return weights, corr_avg


# ── Carga de retornos desde SQLite ────────────────────────────────────────

def _fetch_returns(tickers: list[str], window_days: int,
                   db_path: Path) -> tuple[pd.DataFrame, dict[str, str]]:
    """Carga retornos diarios para tickers. Devuelve (df, skipped_dict)."""
    skipped: dict[str, str] = {}
    if not db_path.exists():
        return pd.DataFrame(), {tk: "db_not_found" for tk in tickers}
    series: dict[str, pd.Series] = {}
    with sqlite3.connect(db_path) as conn:
        for tk in tickers:
            df = pd.read_sql(
                "SELECT date, close FROM prices WHERE ticker = ? ORDER BY date DESC LIMIT ?",
                conn, params=(tk, window_days + 5),
                parse_dates=["date"],
            )
            if df.empty or len(df) < window_days * 0.5:
                skipped[tk] = f"datos_insuficientes_{len(df)}"
                continue
            df = df.sort_values("date").set_index("date")
            r = np.log(df["close"] / df["close"].shift(1)).dropna()
            if len(r) < window_days * 0.5:
                skipped[tk] = "retornos_insuficientes"
                continue
            series[tk] = r.tail(window_days)
    if not series:
        return pd.DataFrame(), skipped
    return pd.DataFrame(series), skipped


def _load_plan_tickers() -> list[str]:
    """Lee paper_report.json y devuelve los tickers del plan que estan en universo Alpaca."""
    p = DASHBOARD_DATA_DIR / "paper_report.json"
    if not p.exists():
        return []
    try:
        data = json.loads(p.read_text(encoding="utf-8"))
        return [t["ticker"] for t in (data.get("plan_today") or [])
                if t.get("in_universe") and t.get("side") == 1]
    except Exception:
        return []


def _load_universe_tickers(db_path: Path) -> list[str]:
    """Lee tickers de active_universe."""
    if not db_path.exists():
        return []
    with sqlite3.connect(db_path) as conn:
        try:
            rows = conn.execute(
                "SELECT ticker FROM active_universe ORDER BY rank ASC NULLS LAST, ticker"
            ).fetchall()
        except sqlite3.OperationalError:
            return []
    return [r[0] for r in rows]


# ── Persistencia ──────────────────────────────────────────────────────────

def _ensure_schema(db_path: Path) -> None:
    with sqlite3.connect(db_path) as conn:
        conn.execute(
            """
            CREATE TABLE IF NOT EXISTS hrp_allocations (
                id INTEGER PRIMARY KEY AUTOINCREMENT,
                timestamp TEXT NOT NULL,
                window_days INTEGER NOT NULL,
                method TEXT NOT NULL,
                tickers TEXT NOT NULL,
                weights TEXT NOT NULL,
                correlation_avg REAL
            )
            """
        )
        conn.commit()


def _persist(report: HrpReport, db_path: Path) -> None:
    with sqlite3.connect(db_path) as conn:
        conn.execute(
            """
            INSERT INTO hrp_allocations
                (timestamp, window_days, method, tickers, weights, correlation_avg)
            VALUES (?, ?, ?, ?, ?, ?)
            """,
            (
                report.generated_at, int(report.window_days), report.method,
                json.dumps(report.tickers_used),
                json.dumps({w.ticker: w.weight for w in report.weights}),
                float(report.correlation_avg) if report.correlation_avg is not None else None,
            ),
        )
        conn.commit()


# ── API publica ───────────────────────────────────────────────────────────

def compute_hrp_for_tickers(
    tickers: list[str],
    window_days: int = DEFAULT_WINDOW,
    db_path: Path = DB_PATH,
) -> dict[str, float]:
    """
    Devuelve dict {ticker: weight} con HRP. Si solo hay 1 ticker, peso=1.0.
    Si los datos son insuficientes para alguno, ese ticker se omite.
    """
    if not tickers:
        return {}
    if len(tickers) == 1:
        return {tickers[0]: 1.0}
    returns, skipped = _fetch_returns(tickers, window_days, db_path)
    if returns.empty or returns.shape[1] < 2:
        # Fallback: equal weight entre los disponibles
        usable = [t for t in tickers if t not in skipped]
        if not usable:
            return {}
        eq = 1.0 / len(usable)
        return {t: eq for t in usable}
    weights, _ = hrp_weights(returns)
    return {tk: float(w) for tk, w in weights.items()}


def annualized_vol(returns: pd.Series) -> float:
    if returns is None or len(returns) < 5:
        return 0.0
    return float(returns.std(ddof=1) * np.sqrt(252))


class AgenteHrp:
    def __init__(self, db_path: Path = DB_PATH, report_path: Path = REPORT_PATH,
                 window: int = DEFAULT_WINDOW) -> None:
        self.db_path = db_path
        self.report_path = report_path
        self.window = window
        self.logger = self._build_logger()
        DASHBOARD_DATA_DIR.mkdir(parents=True, exist_ok=True)
        _ensure_schema(self.db_path)

    def run(self, tickers: list[str], label: str = "manual") -> HrpReport:
        now_iso = datetime.now(UTC).isoformat()
        tickers = list(dict.fromkeys(tickers))  # eliminar duplicados manteniendo orden
        self.logger.info("HRP solicitado para %d tickers (window=%dd, label=%s)",
                         len(tickers), self.window, label)

        if not tickers:
            return self._empty_report(now_iso, [], "No input tickers")

        returns, skipped = _fetch_returns(tickers, self.window, self.db_path)
        if returns.empty:
            self.logger.warning("No data for HRP. Skipped: %s", skipped)
            return self._empty_report(now_iso, tickers, "No data in SQLite", skipped=skipped)

        tickers_used = list(returns.columns)
        if len(tickers_used) < MIN_TICKERS:
            # Caso degenerado: 1 ticker → peso 1.0
            tk = tickers_used[0]
            vol = annualized_vol(returns[tk])
            report = HrpReport(
                generated_at=now_iso, window_days=self.window,
                n_tickers=1, method="single_asset",
                weights=[HrpWeight(ticker=tk, weight=1.0, rank=1, annual_vol=vol)],
                tickers_requested=tickers, tickers_used=tickers_used,
                skipped=skipped, correlation_avg=None,
                note=f"Only {len(tickers_used)} ticker available — weight 1.0",
            )
            _persist(report, self.db_path)
            self._export(report)
            return report

        weights, corr_avg = hrp_weights(returns)
        sorted_w = weights.sort_values(ascending=False)
        weights_out: list[HrpWeight] = []
        for rank, (tk, w) in enumerate(sorted_w.items(), start=1):
            vol = annualized_vol(returns[tk])
            weights_out.append(HrpWeight(
                ticker=tk, weight=round(float(w), 6),
                rank=rank, annual_vol=round(vol, 4),
            ))

        report = HrpReport(
            generated_at=now_iso, window_days=self.window,
            n_tickers=len(tickers_used), method="HRP",
            weights=weights_out, tickers_requested=tickers,
            tickers_used=tickers_used, skipped=skipped,
            correlation_avg=round(corr_avg, 4),
            note=f"HRP over {len(tickers_used)} tickers, avg_corr={corr_avg:.3f}",
        )
        _persist(report, self.db_path)
        self._export(report)
        self.logger.info(
            "HRP OK | tickers=%d | corr_avg=%.3f | max_w=%s (%.2f%%) | min_w=%s (%.2f%%)",
            len(tickers_used), corr_avg,
            weights_out[0].ticker, weights_out[0].weight * 100,
            weights_out[-1].ticker, weights_out[-1].weight * 100,
        )
        return report

    def _empty_report(self, now_iso: str, requested: list[str], note: str,
                      skipped: Optional[dict] = None) -> HrpReport:
        report = HrpReport(
            generated_at=now_iso, window_days=self.window, n_tickers=0,
            method="fallback_equal_weight",
            tickers_requested=requested, tickers_used=[],
            skipped=skipped or {}, note=note,
        )
        self._export(report)
        return report

    def _export(self, report: HrpReport) -> None:
        payload = {
            **{k: v for k, v in asdict(report).items() if k != "weights"},
            "weights": [asdict(w) for w in report.weights],
        }
        self.report_path.write_text(
            json.dumps(payload, indent=2, ensure_ascii=False),
            encoding="utf-8",
        )
        self.logger.info("Informe exportado → %s", self.report_path)

    def _build_logger(self) -> logging.Logger:
        LOGS_DIR.mkdir(parents=True, exist_ok=True)
        logger = logging.getLogger("agente_hrp")
        logger.setLevel(logging.INFO)
        logger.handlers.clear()
        fmt = logging.Formatter("%(asctime)s | %(levelname)s | %(name)s | %(message)s")
        fh = logging.FileHandler(LOGS_DIR / "agente_hrp.log", encoding="utf-8")
        fh.setFormatter(fmt)
        sh = logging.StreamHandler()
        sh.setFormatter(fmt)
        logger.addHandler(fh)
        logger.addHandler(sh)
        return logger


def main() -> int:
    parser = argparse.ArgumentParser(description="Agente HRP de Leonex")
    parser.add_argument("--tickers", type=str, default=None,
                        help="Lista de tickers separados por coma (ej: SPY,QQQ,AAPL)")
    parser.add_argument("--window", type=int, default=DEFAULT_WINDOW,
                        help="Ventana de retornos para covarianza (default 60d)")
    parser.add_argument("--plan", action="store_true",
                        help="Usar tickers del plan diario operables en Alpaca")
    parser.add_argument("--universe", action="store_true",
                        help="Usar todos los tickers del active_universe")
    args = parser.parse_args()

    if sys.stdout.encoding and sys.stdout.encoding.lower() != "utf-8":
        import io
        sys.stdout = io.TextIOWrapper(sys.stdout.buffer, encoding="utf-8", errors="replace")

    # Seleccionar conjunto de tickers
    if args.tickers:
        tickers = [t.strip() for t in args.tickers.split(",") if t.strip()]
        label = "cli_tickers"
    elif args.plan:
        tickers = _load_plan_tickers()
        label = "plan_today"
    elif args.universe:
        tickers = _load_universe_tickers(DB_PATH)
        label = "active_universe"
    else:
        # Default: si hay plan_today, lo usa; si no, active_universe
        tickers = _load_plan_tickers()
        if tickers:
            label = "plan_today_default"
        else:
            tickers = _load_universe_tickers(DB_PATH)
            label = "active_universe_default"

    if not tickers:
        # Plan vacio es un estado NORMAL del mercado (no todos los dias hay
        # senales). NO es un error: devolvemos 0 y un reporte vacio para que
        # el pipeline no se rompa. El dashboard mostrara "sin plan hoy".
        print("[i] No tickers to process today (empty plan or empty universe).")
        print("    This is normal when there are no filtered signals. HRP has no work.")
        try:
            empty_payload = {
                "generated_at": __import__("datetime").datetime.now(
                    __import__("datetime").timezone.utc).isoformat(),
                "method": "none",
                "label": "empty_plan",
                "window_days": args.window,
                "n_tickers": 0,
                "correlation_avg": None,
                "weights": [],
                "note": "No operable tickers today — HRP does not apply.",
            }
            REPORT_PATH.write_text(
                __import__("json").dumps(empty_payload, indent=2, ensure_ascii=False),
                encoding="utf-8",
            )
        except Exception:
            pass
        return 0

    agente = AgenteHrp(window=args.window)
    report = agente.run(tickers, label=label)

    sep = "=" * 64
    print(f"\n{sep}")
    print("Leonex -- HRP (Hierarchical Risk Parity)")
    print(sep)
    print(f"Metodo            : {report.method}")
    print(f"Ventana           : {report.window_days} dias")
    print(f"Tickers usados    : {report.n_tickers}")
    if report.correlation_avg is not None:
        print(f"Corr media        : {report.correlation_avg:+.3f}")
    print(f"Nota              : {report.note}")
    if report.weights:
        print()
        print(f"  {'rank':>4}  {'ticker':<10} {'weight':>9}  {'ann_vol':>8}  {'$ por $10K':>11}")
        for w in report.weights:
            ten_k = w.weight * 10000
            print(f"  {w.rank:>4}  {w.ticker:<10} {w.weight*100:>7.2f}%  "
                  f"{w.annual_vol*100:>6.2f}%  ${ten_k:>10.2f}")
    if report.skipped:
        print()
        print(f"Skipped ({len(report.skipped)}):")
        for tk, reason in list(report.skipped.items())[:10]:
            print(f"  {tk:<10} {reason}")
    print(sep)
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
