#!/usr/bin/env python3
"""
**Johansen triplet stat-arb** on S&P 500 (Chan Ex 2.7–2.8), sector-aware scan.

Cannot exhaust C(500,3) triplets; we scan **within GICS sectors** on a training
window, rank by Johansen trace + half-life, greedily pick non-overlapping books,
then equal-weight the triplet sleeves.

Discovery (training window, default 252 sessions ending at ``--start``):
  1. Load S&P 500 tickers + GICS sector (Wikipedia or ``alpaca_symbols.txt`` fallback)
  2. Per sector: combinations of 3 among up to ``--max-per-sector`` names
  3. Johansen trace rejects r=0 at 95%; half-life in ``[hl_min, hl_max]`` days
  4. Greedy select ``--top-k`` triplets (no ticker overlap)

Trading: causal Johansen refit (default) on each selected triplet; portfolio = mean of sleeves.

Example (smoke — 80 names, 3 triplets)::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 \\
      .venv/bin/python RenTech/strategy_stack/run_johansen_triplet_sp500.py \\
      --start 2018-01-04 --end 2026-04-02 --capital 100000 \\
      --max-universe 80 --top-k 3 --train-days 252

Full universe (slow scan + download)::

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

from __future__ import annotations

import argparse
import io
import itertools
import json
import sys
import urllib.request
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,
    johansen_trace_pvalue,
    unit_portfolio_half_life,
)
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_sp500"
_WIKI_SP500 = "https://en.wikipedia.org/wiki/List_of_S%26P_500_companies"
_USER_AGENT = (
    "Mozilla/5.0 (compatible; trading_bot/1.0) AppleWebKit/537.36 "
    "(KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36"
)


@dataclass(frozen=True)
class TripletCandidate:
    tickers: tuple[str, str, str]
    sector: str
    trace_stat: float
    crit_95: float
    half_life: float
    score: float


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 load_sp500_sectors(*, symbols_file: Path | None = None) -> pd.DataFrame:
    """
    Return DataFrame ``ticker``, ``sector`` (GICS).

    Falls back to flat list from ``alpaca_symbols.txt`` with sector ``Unknown``.
    """
    try:
        req = urllib.request.Request(_WIKI_SP500, headers={"User-Agent": _USER_AGENT})
        with urllib.request.urlopen(req, timeout=30) as resp:
            html = resp.read().decode("utf-8", errors="replace")
        tables = pd.read_html(io.StringIO(html))
        df = tables[0]
        sym_col = "Symbol" if "Symbol" in df.columns else "Ticker"
        sec_col = next(
            (c for c in df.columns if str(c).strip().lower() in ("gics sector", "sector")),
            None,
        )
        if sym_col not in df.columns:
            raise KeyError("no symbol column")
        out = pd.DataFrame(
            {
                "ticker": df[sym_col].astype(str).str.strip().str.upper().str.replace(".", "-", regex=False),
                "sector": df[sec_col].astype(str) if sec_col else "Unknown",
            }
        )
        return out.drop_duplicates("ticker").reset_index(drop=True)
    except Exception as exc:
        path = symbols_file or (_REPO / "alpaca_symbols.txt")
        if not path.is_file():
            raise RuntimeError(f"Wikipedia SP500 scrape failed ({exc}) and no {path}") from exc
        tickers = [ln.strip().upper() for ln in path.read_text().splitlines() if ln.strip()]
        return pd.DataFrame({"ticker": tickers, "sector": "Unknown"})


def _greedy_select_triplets(
    candidates: list[TripletCandidate],
    top_k: int,
) -> list[TripletCandidate]:
    """Highest score first; no ticker may appear in more than one triplet."""
    used: set[str] = set()
    picked: list[TripletCandidate] = []
    for c in sorted(candidates, key=lambda x: x.score, reverse=True):
        legs = set(c.tickers)
        if legs & used:
            continue
        picked.append(c)
        used |= legs
        if len(picked) >= top_k:
            break
    return picked


def scan_sector_triplets(
    close: pd.DataFrame,
    tickers: list[str],
    sector: str,
    train_mask: pd.Series,
    *,
    min_train: int,
    half_life_min: float,
    half_life_max: float,
    max_combos: int,
) -> list[TripletCandidate]:
    """Enumerate 3-combos in sector on training slice; Johansen + half-life filter."""
    avail = [t for t in tickers if t in close.columns]
    if len(avail) < 3:
        return []
    train = close.loc[train_mask, avail].dropna(how="any")
    if len(train) < min_train:
        return []

    combos = list(itertools.combinations(sorted(avail), 3))
    if len(combos) > max_combos:
        rng = np.random.default_rng(42)
        idx = rng.choice(len(combos), size=max_combos, replace=False)
        combos = [combos[i] for i in sorted(idx)]

    out: list[TripletCandidate] = []
    for legs in combos:
        px = train[list(legs)]
        if len(px) < min_train:
            continue
        try:
            trace, crit95, evec = johansen_trace_pvalue(px)
        except Exception:
            continue
        if trace <= crit95:
            continue
        hl = unit_portfolio_half_life(px, evec)
        if not np.isfinite(hl) or hl < half_life_min or hl > half_life_max:
            continue
        # Prefer strong cointegration + fast mean reversion
        score = float(trace / max(crit95, 1e-9) - 0.05 * hl)
        out.append(
            TripletCandidate(
                tickers=legs,
                sector=sector,
                trace_stat=trace,
                crit_95=crit95,
                half_life=float(hl),
                score=score,
            )
        )
    return out


def discover_triplets(
    close: pd.DataFrame,
    sectors: pd.DataFrame,
    *,
    train_end: pd.Timestamp,
    train_days: int,
    max_per_sector: int,
    max_combos_per_sector: int,
    min_train: int,
    half_life_min: float,
    half_life_max: float,
) -> list[TripletCandidate]:
    train_end = pd.Timestamp(train_end).normalize()
    train_start = train_end - pd.Timedelta(days=int(train_days * 1.6))
    train_mask = (close.index >= train_start) & (close.index <= train_end)

    # Cap per sector by training-window average price level (liquidity proxy)
    candidates: list[TripletCandidate] = []
    for sector, grp in sectors.groupby("sector"):
        tickers = [t for t in grp["ticker"].tolist() if t in close.columns]
        if len(tickers) < 3:
            continue
        if len(tickers) > max_per_sector:
            avg_px = close.loc[train_mask, tickers].mean().sort_values(ascending=False)
            tickers = list(avg_px.head(max_per_sector).index)
        found = scan_sector_triplets(
            close,
            tickers,
            str(sector),
            train_mask,
            min_train=min_train,
            half_life_min=half_life_min,
            half_life_max=half_life_max,
            max_combos=max_combos_per_sector,
        )
        candidates.extend(found)
    return candidates


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"), "sharpe": 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 _backtest_triplet_chunk(
    close: pd.DataFrame,
    legs: tuple[str, ...],
    *,
    chunk_start: pd.Timestamp,
    chunk_end: pd.Timestamp,
    fidelity: Fidelity,
    refit_bars: int,
    min_train: int,
) -> pd.Series:
    px = close[list(legs)].dropna(how="any")
    r, _ = johansen_linear_triplet(
        px, fidelity=fidelity, refit_bars=refit_bars, min_train=min_train
    )
    r.index = pd.to_datetime(r.index).tz_localize(None)
    mask = (r.index >= chunk_start) & (r.index <= chunk_end)
    return r.loc[mask].fillna(0.0)


def run_walk_forward(
    close: pd.DataFrame,
    sectors_df: pd.DataFrame,
    *,
    t0: pd.Timestamp,
    t1: pd.Timestamp,
    walk_months: int,
    train_days: int,
    top_k: int,
    fidelity: Fidelity,
    refit_bars: int,
    min_train: int,
    scan_kw: dict[str, Any],
) -> tuple[pd.Series, list[dict[str, Any]]]:
    """Rescan triplets every ``walk_months``; trade only the following window."""
    anchors = pd.date_range(t0, t1, freq=f"{int(walk_months)}MS")
    if len(anchors) < 2:
        anchors = pd.DatetimeIndex([t0, t1])
    sleeve_chunks: list[pd.Series] = []
    wf_log: list[dict[str, Any]] = []
    for i in range(len(anchors) - 1):
        chunk_start = pd.Timestamp(anchors[i]).normalize()
        chunk_end = (pd.Timestamp(anchors[i + 1]) - pd.Timedelta(days=1)).normalize()
        if chunk_end < chunk_start:
            continue
        train_end = chunk_start - pd.Timedelta(days=1)
        cands = discover_triplets(
            close,
            sectors_df,
            train_end=train_end,
            train_days=train_days,
            **scan_kw,
        )
        picked = _greedy_select_triplets(cands, top_k)
        if not picked:
            continue
        chunk_rets: list[pd.Series] = []
        for c in picked:
            try:
                chunk_rets.append(
                    _backtest_triplet_chunk(
                        close,
                        c.tickers,
                        chunk_start=chunk_start,
                        chunk_end=chunk_end,
                        fidelity=fidelity,
                        refit_bars=refit_bars,
                        min_train=min_train,
                    )
                )
            except Exception:
                continue
        if not chunk_rets:
            continue
        chunk_port = pd.concat(chunk_rets, axis=1).fillna(0.0).mean(axis=1)
        sleeve_chunks.append(chunk_port)
        wf_log.append(
            {
                "chunk_start": str(chunk_start.date()),
                "chunk_end": str(chunk_end.date()),
                "triplets": ["-".join(c.tickers) for c in picked],
            }
        )
    if not sleeve_chunks:
        return pd.Series(dtype=float), wf_log
    port = pd.concat(sleeve_chunks).sort_index()
    port = port[~port.index.duplicated(keep="last")]
    return port.fillna(0.0), wf_log


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
    ap.add_argument("--start", default="2016-01-04", help="Backtest start (also train end for discovery)")
    ap.add_argument("--end", default="2026-04-02")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--top-k", type=int, default=5, help="Max non-overlapping triplets in book")
    ap.add_argument("--train-days", type=int, default=252, help="Training window length (sessions)")
    ap.add_argument("--min-train", type=int, default=252, help="Min bars for Johansen inside backtest")
    ap.add_argument("--lookback-years", type=float, default=12.0)
    ap.add_argument("--max-universe", type=int, default=0, help="Cap SP500 download count (0=all)")
    ap.add_argument("--max-per-sector", type=int, default=25, help="Max names per sector for triplet scan")
    ap.add_argument(
        "--max-combos-per-sector",
        type=int,
        default=800,
        help="Cap random triplet combos per sector (speed)",
    )
    ap.add_argument("--half-life-min", type=float, default=1.0)
    ap.add_argument("--half-life-max", type=float, default=40.0)
    ap.add_argument(
        "--fidelity",
        choices=("book", "causal"),
        default="causal",
        help="causal recommended for OOS stock triplets",
    )
    ap.add_argument("--refit-bars", type=int, default=63)
    ap.add_argument("--refresh-cache", action="store_true")
    ap.add_argument("--symbols-file", type=Path, default=_REPO / "alpaca_symbols.txt")
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT_PREFIX)
    ap.add_argument(
        "--triplets",
        default="",
        help="Skip scan; fixed triplet e.g. AAPL,MSFT,GOOGL|XOM,CVX,COP",
    )
    ap.add_argument(
        "--walk-forward-months",
        type=int,
        default=0,
        help="If >0, rescan triplets every N months on trailing train-days (OOS discipline)",
    )
    args = ap.parse_args()

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

    print("Johansen triplet SP500 book", flush=True)
    sectors_df = load_sp500_sectors(symbols_file=args.symbols_file)
    if int(args.max_universe) > 0:
        cap_n = int(args.max_universe)
        # Stratified sample — avoid alphabetical bias from raw list order
        parts: list[pd.DataFrame] = []
        for _, grp in sectors_df.groupby("sector"):
            n_take = max(1, int(round(cap_n * len(grp) / len(sectors_df))))
            parts.append(grp.sample(n=min(n_take, len(grp)), random_state=42))
        sectors_df = pd.concat(parts).drop_duplicates("ticker").head(cap_n)
    tickers = sectors_df["ticker"].tolist()

    print(f"  universe: {len(tickers)} names  train_days={args.train_days}  top_k={args.top_k}", 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)

    scan_kw = dict(
        max_per_sector=int(args.max_per_sector),
        max_combos_per_sector=int(args.max_combos_per_sector),
        min_train=min(int(args.train_days), 200),
        half_life_min=float(args.half_life_min),
        half_life_max=float(args.half_life_max),
    )

    wf_months = int(args.walk_forward_months)
    if wf_months > 0 and not str(args.triplets).strip():
        print(f"  walk-forward: rescan every {wf_months} months", flush=True)
        port_r, wf_log = run_walk_forward(
            close,
            sectors_df,
            t0=t0,
            t1=t1 or pd.Timestamp(close.index.max()),
            walk_months=wf_months,
            train_days=int(args.train_days),
            top_k=int(args.top_k),
            fidelity=fidelity,
            refit_bars=int(args.refit_bars),
            min_train=int(args.min_train),
            scan_kw=scan_kw,
        )
        if port_r.empty:
            raise SystemExit("Walk-forward produced no returns.")
        picked = []  # logged in wf_log
        sleeve_meta = [{"walk_forward": wf_log}]
        pm = _metrics(port_r, capital=cap)
        print(
            f"\nPORTFOLIO (walk-forward)  ret={pm['total_return_pct']:.1f}%  "
            f"CAGR={pm['cagr_pct']:.1f}%  Sharpe={pm['sharpe']:.2f}  maxDD={pm['max_drawdown_pct']:.1f}%",
            flush=True,
        )
        prefix = args.out_prefix.expanduser().resolve()
        prefix.parent.mkdir(parents=True, exist_ok=True)
        eq = cap * (1.0 + port_r).cumprod()
        pd.DataFrame(
            {
                "date": port_r.index.strftime("%Y-%m-%d"),
                "daily_ret": port_r.values,
                "daily_pnl_usd": (port_r * cap).values,
                "equity_usd": eq.values,
            }
        ).to_csv(f"{prefix}_daily.csv", index=False)
        Path(f"{prefix}_meta.json").write_text(
            json.dumps(
                {
                    "strategy": "chan_johansen_triplet_sp500_walk_forward",
                    "command": " ".join(sys.argv),
                    "portfolio": pm,
                    "walk_forward_log": wf_log,
                },
                indent=2,
            )
        )
        print(f"\nWrote {prefix}_daily.csv", flush=True)
        print(f"Wrote {prefix}_meta.json", flush=True)
        return

    if str(args.triplets).strip():
        picked: list[TripletCandidate] = []
        for chunk in str(args.triplets).split("|"):
            legs = tuple(x.strip().upper() for x in chunk.split(",") if x.strip())
            if len(legs) != 3:
                raise SystemExit(f"--triplets chunk must have 3 tickers: {chunk!r}")
            picked.append(
                TripletCandidate(
                    tickers=legs,  # type: ignore[arg-type]
                    sector="fixed",
                    trace_stat=float("nan"),
                    crit_95=float("nan"),
                    half_life=float("nan"),
                    score=0.0,
                )
            )
    else:
        print("  scanning sector triplets on training window …", flush=True)
        all_cand = discover_triplets(
            close,
            sectors_df,
            train_end=t0,
            train_days=int(args.train_days),
            **scan_kw,
        )
        print(f"  raw candidates: {len(all_cand)}", flush=True)
        picked = _greedy_select_triplets(all_cand, int(args.top_k))
        if not picked:
            raise SystemExit(
                "No triplets passed filters. Try --max-universe 150, widen --half-life-max, "
                "or pass --triplets AAPL,MSFT,GOOGL"
            )

    print("  selected triplets:", flush=True)
    for i, c in enumerate(picked, 1):
        print(
            f"    {i}. {'-'.join(c.tickers)}  sector={c.sector}  "
            f"trace={c.trace_stat:.1f}  hl={c.half_life:.1f}d  score={c.score:.2f}",
            flush=True,
        )

    sleeve_rets: dict[str, pd.Series] = {}
    sleeve_meta: list[dict[str, Any]] = []
    for c in picked:
        legs = list(c.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=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,
                    "sector": c.sector,
                    "trace_stat": c.trace_stat,
                    "half_life_train": c.half_life,
                    **m,
                    **meta,
                }
            )
            print(
                f"  {'-'.join(legs):22s}  ret={m['total_return_pct']:6.1f}%  "
                f"Sharpe={m['sharpe']:5.2f}  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)  ret={pm['total_return_pct']:.1f}%  "
        f"CAGR={pm['cagr_pct']:.1f}%  Sharpe={pm['sharpe']:.2f}  maxDD={pm['max_drawdown_pct']:.1f}%",
        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_path = Path(f"{prefix}_meta.json")
    meta_path.write_text(
        json.dumps(
            {
                "strategy": "chan_johansen_triplet_sp500",
                "command": " ".join(sys.argv),
                "start": args.start,
                "end": args.end,
                "capital_usd": cap,
                "fidelity": fidelity,
                "portfolio": pm,
                "beta_spy": beta,
                "rho_spy": rho,
                "triplets": sleeve_meta,
                "caveats": [
                    "Survivorship bias: current SP500 list used for full history",
                    "No transaction costs or short-sale constraints",
                    "Sector scan is heuristic; not exhaustive C(500,3)",
                ],
            },
            indent=2,
            default=str,
        )
    )
    print(f"\nWrote {prefix}_daily.csv", flush=True)
    print(f"Wrote {prefix}_triplets.csv", flush=True)
    print(f"Wrote {meta_path}", flush=True)


if __name__ == "__main__":
    main()
