#!/usr/bin/env python3
"""
Equal-weight **Johansen triplet** portfolio on curated ETF books.

Default book: top-5 OOS triplets from ``run_johansen_triplet_scan.py`` plus
Chan's **EWA–EWC–IGE** classic (book-fidelity expanding Johansen on that sleeve).

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 \\
      .venv/bin/python RenTech/strategy_stack/run_johansen_triplet_etf_portfolio.py \\
      --start 2016-01-04 --end 2026-04-02 --capital 100000
"""

from __future__ import annotations

import argparse
import json
import sys
from dataclasses import dataclass
from pathlib import Path
from typing import Any

import numpy as np
import pandas as pd

_REPO = Path(__file__).resolve().parents[2]
if str(_REPO) not in sys.path:
    sys.path.insert(0, str(_REPO))

import universe_scanner as us  # type: ignore[import-not-found]

from RenTech.strategy_stack.chan_johansen_triplet import Fidelity, johansen_linear_triplet
from RenTech.strategy_stack.data_loader import DataLoader
from RenTech.strategy_stack.main import _compute_daily_backtest_features

LOGS = _REPO / "RenTech" / "data" / "logs"
DEFAULT_OUT_PREFIX = LOGS / "johansen_triplet_etf_standard"

# (legs, label, per-sleeve fidelity override or None → --fidelity)
DEFAULT_SLEEVES: list[tuple[tuple[str, str, str], str, Fidelity | None]] = [
    (("GDXJ", "IAU", "SIL"), "precious_metals_junior", None),
    (("GLD", "UNG", "USO"), "commodity_gold_gas_oil", None),
    (("XLB", "XLI", "XLP"), "sector_cyclical_defensive", None),
    (("COP", "USO", "XOP"), "energy_complex", None),
    (("DBC", "PDBC", "USO"), "commodity_broad", None),
    (("EWA", "EWC", "IGE"), "chan_classic", "book"),
]


@dataclass(frozen=True)
class SleeveSpec:
    tickers: tuple[str, str, str]
    label: str
    fidelity: Fidelity


def _close_panel(data: dict[str, pd.DataFrame]) -> pd.DataFrame:
    cols: dict[str, pd.Series] = {}
    for t, df in data.items():
        if df is None or df.empty:
            continue
        c = "close" if "close" in df.columns else "Close"
        if c not in df.columns:
            continue
        s = df[c].astype(np.float64)
        s.index = pd.to_datetime(s.index).tz_localize(None)
        cols[t] = s
    if not cols:
        raise ValueError("empty close panel")
    return pd.DataFrame(cols).sort_index()


def _metrics(r: pd.Series, *, capital: float) -> dict[str, float]:
    r = r.fillna(0.0).astype(np.float64)
    n = len(r)
    if n < 2:
        return {
            "n_days": float(n),
            "total_return_pct": float("nan"),
            "cagr_pct": float("nan"),
            "sharpe": float("nan"),
            "max_drawdown_pct": float("nan"),
            "ending_equity_usd": float("nan"),
        }
    eq = capital * (1.0 + r).cumprod()
    years = n / 252.0
    tot = float(eq.iloc[-1] / capital - 1.0)
    cagr = (eq.iloc[-1] / capital) ** (1.0 / years) - 1.0 if years > 0 else float("nan")
    sd = float(r.std(ddof=1))
    sharpe = float(r.mean() / sd * np.sqrt(252.0)) if sd > 1e-12 else float("nan")
    mdd = float((eq / eq.cummax() - 1.0).min())
    return {
        "n_days": float(n),
        "total_return_pct": tot * 100.0,
        "cagr_pct": float(cagr * 100.0),
        "sharpe": sharpe,
        "max_drawdown_pct": mdd * 100.0,
        "ending_equity_usd": float(eq.iloc[-1]),
    }


def _parse_sleeves(
    triplet_arg: str,
    *,
    default_fidelity: Fidelity,
) -> list[SleeveSpec]:
    if not str(triplet_arg).strip():
        return [
            SleeveSpec(tickers=legs, label=label, fidelity=fid or default_fidelity)
            for legs, label, fid in DEFAULT_SLEEVES
        ]
    out: list[SleeveSpec] = []
    for chunk in str(triplet_arg).split("|"):
        parts = [x.strip() for x in chunk.split(",") if x.strip()]
        if len(parts) != 3:
            raise SystemExit(f"--triplets chunk must have 3 tickers: {chunk!r}")
        legs = (parts[0].upper(), parts[1].upper(), parts[2].upper())
        out.append(SleeveSpec(tickers=legs, label="custom", fidelity=default_fidelity))
    return out


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
    ap.add_argument("--start", default="2016-01-04")
    ap.add_argument("--end", default="2026-04-02")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--min-train", type=int, default=252)
    ap.add_argument("--refit-bars", type=int, default=63)
    ap.add_argument(
        "--fidelity",
        choices=("book", "causal"),
        default="causal",
        help="Default Johansen mode; chan classic sleeve uses book unless overridden in code",
    )
    ap.add_argument("--lookback-years", type=float, default=15.0)
    ap.add_argument("--refresh-cache", action="store_true")
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT_PREFIX)
    ap.add_argument(
        "--triplets",
        default="",
        help="Override default book: A,B,C|D,E,F (pipe-separated)",
    )
    args = ap.parse_args()

    t0 = pd.Timestamp(args.start).normalize()
    t1 = pd.Timestamp(args.end).normalize() if str(args.end).strip() else None
    cap = float(args.capital)
    default_fidelity: Fidelity = str(args.fidelity)  # type: ignore[assignment]
    sleeves = _parse_sleeves(args.triplets, default_fidelity=default_fidelity)

    tickers = sorted({t for s in sleeves for t in s.tickers})
    print(
        f"Johansen ETF triplet portfolio  {len(sleeves)} sleeves  "
        f"default_fidelity={default_fidelity}",
        flush=True,
    )
    print(f"  window: {t0.date()} → {t1.date() if t1 else 'max'}  capital=${cap:,.0f}", flush=True)

    raw = us.download_and_cache_data(
        tickers,
        timeframe="1d",
        lookback_years=float(args.lookback_years),
        max_age_hours=0.0 if args.refresh_cache else 24.0,
    )
    close = _close_panel(raw)
    print(f"  panel: {close.shape[1]} tickers × {len(close)} days", flush=True)

    sleeve_rets: dict[str, pd.Series] = {}
    sleeve_meta: list[dict[str, Any]] = []

    print("\n  sleeves:", flush=True)
    for spec in sleeves:
        legs = list(spec.tickers)
        missing = [t for t in legs if t not in close.columns]
        if missing:
            print(f"    Skip {'-'.join(legs)}: missing {missing}", flush=True)
            continue
        px = close[legs].dropna(how="any")
        try:
            r, meta = johansen_linear_triplet(
                px,
                fidelity=spec.fidelity,
                refit_bars=int(args.refit_bars),
                min_train=int(args.min_train),
            )
            r = r.copy()
            r.index = pd.to_datetime(r.index).tz_localize(None)
            mask = r.index >= t0
            if t1 is not None:
                mask &= r.index <= t1
            r = r.loc[mask].fillna(0.0)
            slug = "_".join(legs)
            sleeve_rets[slug] = r
            m = _metrics(r, capital=cap)
            sleeve_meta.append(
                {
                    "triplet": legs,
                    "label": spec.label,
                    "fidelity": spec.fidelity,
                    **m,
                    **meta,
                }
            )
            print(
                f"    {'-'.join(legs):22s}  {spec.label:26s}  fid={spec.fidelity:6s}  "
                f"ret={m['total_return_pct']:6.1f}%  Sharpe={m['sharpe']:5.2f}  "
                f"maxDD={m['max_drawdown_pct']:5.1f}%",
                flush=True,
            )
        except Exception as exc:
            print(f"    Skip {'-'.join(legs)}: {exc}", flush=True)

    if not sleeve_rets:
        raise SystemExit("No triplet backtests succeeded.")

    panel = pd.DataFrame(sleeve_rets).fillna(0.0)
    port_r = panel.mean(axis=1)
    pm = _metrics(port_r, capital=cap)
    print(
        f"\nPORTFOLIO ({len(sleeve_rets)} sleeves, equal-weight)  "
        f"ret={pm['total_return_pct']:.1f}%  CAGR={pm['cagr_pct']:.1f}%  "
        f"Sharpe={pm['sharpe']:.2f}  maxDD={pm['max_drawdown_pct']:.1f}%  "
        f"end=${pm['ending_equity_usd']:,.0f}",
        flush=True,
    )

    prefix = args.out_prefix.expanduser().resolve()
    prefix.parent.mkdir(parents=True, exist_ok=True)
    eq = cap * (1.0 + port_r).cumprod()
    pnl = port_r * cap

    pd.DataFrame(
        {
            "date": port_r.index.strftime("%Y-%m-%d"),
            "daily_ret": port_r.values,
            "daily_pnl_usd": pnl.values,
            "equity_usd": eq.values,
            **{f"ret_{c}": panel[c].values for c in panel.columns},
        }
    ).to_csv(f"{prefix}_daily.csv", index=False)

    pd.DataFrame(sleeve_meta).to_csv(f"{prefix}_triplets.csv", index=False)

    spy = _compute_daily_backtest_features(DataLoader().fetch_daily("SPY", period="max"))
    spy_r = spy["ret"].astype(float)
    spy_r.index = pd.to_datetime(spy_r.index).tz_localize(None)
    aligned = pd.DataFrame({"p": port_r, "spy": spy_r.reindex(port_r.index)}).dropna()
    beta = float("nan")
    rho = float("nan")
    if len(aligned) > 20:
        rho = float(aligned.corr().iloc[0, 1])
        sv = float(aligned["spy"].var())
        beta = float(aligned["p"].cov(aligned["spy"]) / sv) if sv > 1e-14 else float("nan")

    meta = {
        "strategy": "chan_johansen_triplet_etf_portfolio",
        "command": " ".join(sys.argv),
        "start": args.start,
        "end": args.end,
        "capital_usd": cap,
        "default_fidelity": default_fidelity,
        "portfolio": pm,
        "beta_spy": beta,
        "rho_spy": rho,
        "sleeves": sleeve_meta,
        "caveats": [
            "USO appears in multiple sleeves; not deduplicated by design",
            "No transaction costs or short-sale constraints",
            "Chan classic EWA-EWC-IGE uses book-fidelity expanding Johansen",
            "Other sleeves use causal Johansen with quarterly refit (refit_bars=63)",
        ],
    }
    meta_path = Path(f"{prefix}_meta.json")
    meta_path.write_text(json.dumps(meta, indent=2, default=str))

    metrics_path = Path(f"{prefix}_metrics.txt")
    metrics_path.write_text(
        "\n".join(
            [
                "# Johansen ETF triplet portfolio",
                f"command: {' '.join(sys.argv)}",
                f"window: {args.start} → {args.end}",
                f"capital: ${cap:,.0f}",
                f"sleeves: {len(sleeve_rets)} equal-weight",
                "",
                f"total_return_pct: {pm['total_return_pct']:.2f}",
                f"cagr_pct: {pm['cagr_pct']:.2f}",
                f"sharpe: {pm['sharpe']:.3f}",
                f"max_drawdown_pct: {pm['max_drawdown_pct']:.2f}",
                f"ending_equity_usd: {pm['ending_equity_usd']:.2f}",
                f"beta_spy: {beta:.3f}",
                f"rho_spy: {rho:.3f}",
                "",
                f"daily_csv: {prefix}_daily.csv",
                f"triplets_csv: {prefix}_triplets.csv",
            ]
        )
        + "\n"
    )

    print(f"\nWrote {prefix}_daily.csv", flush=True)
    print(f"Wrote {prefix}_triplets.csv", flush=True)
    print(f"Wrote {meta_path}", flush=True)
    print(f"Wrote {metrics_path}", flush=True)


if __name__ == "__main__":
    main()
