#!/usr/bin/env python3
"""
Compare buy-the-dip variants on S&P 500 (2016+):

1. **baseline** — −3%% day, rank ATR/px (current ``run_sp500_dip_standard``)
2. **relative** — −3%% and underperformed SPY by ≥2%% same day, rank by underperformance
3. **hedged** — relative signals + short SPY overlay (excess-return sleeve)

Example::

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

from __future__ import annotations

import argparse
import json
import math
import sys
from dataclasses import dataclass
from pathlib import Path

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))

from RenTech.strategy_stack.data_loader import DataLoader
from RenTech.strategy_stack.equity_universe_loaders import load_equity_panel_dict
from RenTech.strategy_stack.main import _compute_daily_backtest_features
from RenTech.strategy_stack.multi_strategy_manager import BuyTheDipSleeve
from RenTech.strategy_stack.run_sp500_dip_standard import (
    DEFAULT_TOP_N_BY_UNIVERSE,
    _filter_equity_by_history,
)

LOGS = _REPO / "RenTech" / "data" / "logs"
DEFAULT_OUT = LOGS / "relative_dip_compare"


@dataclass(frozen=True)
class DipVariantSpec:
    key: str
    label: str
    relative_spy_min: float
    rank_by: str
    hedge_spy: bool


VARIANTS: tuple[DipVariantSpec, ...] = (
    DipVariantSpec(
        "baseline",
        "Baseline (−3%%, rank ATR)",
        relative_spy_min=0.0,
        rank_by="atr_norm",
        hedge_spy=False,
    ),
    DipVariantSpec(
        "relative",
        "Relative dip (−3%% & underperf SPY ≥2%%, rank underperf)",
        relative_spy_min=0.02,
        rank_by="rel_underperf",
        hedge_spy=False,
    ),
    DipVariantSpec(
        "hedged",
        "Relative + SPY hedge (excess return)",
        relative_spy_min=0.02,
        rank_by="rel_underperf",
        hedge_spy=True,
    ),
)


def _metrics(r: pd.Series, spy_r: pd.Series, cap: float) -> dict:
    r = r.fillna(0.0).astype(np.float64)
    n = len(r)
    eq = cap * (1.0 + r).cumprod()
    years = n / 252.0
    end = float(eq.iloc[-1])
    tot = end / cap - 1.0
    cagr = (end / cap) ** (1.0 / years) - 1.0 if years > 0 else float("nan")
    sd = float(r.std(ddof=1)) if n > 1 else float("nan")
    sharpe = float(r.mean() / sd * math.sqrt(252.0)) if sd > 1e-12 else 0.0
    dd = float((eq / eq.cummax() - 1.0).min())

    aligned = pd.DataFrame({"dip": r, "spy": spy_r}).dropna()
    rho = float(aligned.corr().iloc[0, 1]) if len(aligned) > 2 else float("nan")
    beta = float("nan")
    if len(aligned) > 10:
        spy_v = float(aligned["spy"].var())
        if spy_v > 1e-14:
            beta = float(aligned["dip"].cov(aligned["spy"]) / spy_v)

    yearly: dict[str, float] = {}
    for yr in sorted(r.index.year.unique()):
        yr_mask = r.index.year == yr
        yr_eq = eq.loc[yr_mask]
        if yr_mask.any():
            start_idx = r.index.get_indexer([yr_eq.index[0]])[0]
            prev_eq = cap if start_idx == 0 else float(eq.iloc[start_idx - 1])
            yearly[str(int(yr))] = round((float(yr_eq.iloc[-1]) / prev_eq - 1.0) * 100.0, 2)

    return {
        "n_sessions": n,
        "total_return_pct": round(tot * 100.0, 2),
        "cagr_pct": round(cagr * 100.0, 2),
        "sharpe": round(sharpe, 3),
        "max_drawdown_pct": round(dd * 100.0, 2),
        "vol_ann_pct": round(sd * math.sqrt(252.0) * 100.0, 2) if math.isfinite(sd) else float("nan"),
        "corr_vs_spy": round(rho, 3),
        "beta_vs_spy": round(beta, 3),
        "ending_equity_usd": round(end, 0),
        "yearly_return_pct_chained": yearly,
    }


def _run_variant(
    spec: DipVariantSpec,
    *,
    equity_dict: dict,
    spy_df: pd.DataFrame,
    top_n: int,
    hold_days: int,
    pct_drop: float,
    start: str,
    end: str,
    cap: float,
) -> tuple[pd.Series, dict]:
    eng = BuyTheDipSleeve(
        signal_mode="pct_drop",
        pct_drop_min=pct_drop,
        hold_trading_days=hold_days,
        relative_spy_min=spec.relative_spy_min,
        rank_by=spec.rank_by,
        hedge_spy=spec.hedge_spy,
    )
    daily_ret = eng.generate_returns(
        equity_dict,
        top_n=top_n,
        spy_df=spy_df,
        verbose=True,
    )
    daily_ret = daily_ret.sort_index()
    daily_ret.index = pd.to_datetime(daily_ret.index).tz_localize(None)
    mask = daily_ret.index >= pd.Timestamp(start)
    if end.strip():
        mask &= daily_ret.index <= pd.Timestamp(end)
    r = daily_ret.loc[mask].fillna(0.0).astype(np.float64)

    spy_r = spy_df["ret"].astype(float)
    spy_r.index = pd.to_datetime(spy_r.index).tz_localize(None)
    spy_r = spy_r.reindex(r.index).fillna(0.0)

    meta = _metrics(r, spy_r, cap)
    meta.update(
        {
            "variant": spec.key,
            "label": spec.label,
            "relative_spy_min": spec.relative_spy_min,
            "rank_by": spec.rank_by,
            "hedge_spy": spec.hedge_spy,
            "top_n": top_n,
            "hold_days": hold_days,
            "pct_drop_min": pct_drop,
            "start": str(r.index.min().date()),
            "end": str(r.index.max().date()),
        }
    )
    return r, meta


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("--yahoo-period", default="max")
    ap.add_argument("--top-n", type=int, default=DEFAULT_TOP_N_BY_UNIVERSE["sp500"])
    ap.add_argument("--hold-days", type=int, default=10)
    ap.add_argument("--pct-drop-min", type=float, default=0.03)
    ap.add_argument("--relative-spy-min", type=float, default=0.02, help="For relative/hedged variants")
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT)
    ap.add_argument("--refresh-cache", action="store_true")
    args = ap.parse_args()

    min_first = pd.Timestamp(args.start).normalize() - pd.Timedelta(days=400)
    equity_dict = load_equity_panel_dict(
        "sp500",
        args.yahoo_period,
        refresh_cache=bool(args.refresh_cache),
    )
    n_loaded = len(equity_dict)
    equity_dict = _filter_equity_by_history(equity_dict, min_first)
    if len(equity_dict) < 50:
        raise SystemExit(f"Too few sp500 names ({len(equity_dict)} of {n_loaded})")

    spy_df = _compute_daily_backtest_features(
        DataLoader().fetch_daily("SPY", period=args.yahoo_period)
    )

    variants = list(VARIANTS)
    # Allow CLI override of relative threshold on non-baseline variants.
    rel_min = float(args.relative_spy_min)
    variants[1] = DipVariantSpec(
        variants[1].key,
        variants[1].label,
        relative_spy_min=rel_min,
        rank_by=variants[1].rank_by,
        hedge_spy=variants[1].hedge_spy,
    )
    variants[2] = DipVariantSpec(
        variants[2].key,
        variants[2].label,
        relative_spy_min=rel_min,
        rank_by=variants[2].rank_by,
        hedge_spy=variants[2].hedge_spy,
    )

    cap = float(args.capital)
    prefix = args.out_prefix.expanduser().resolve()
    prefix.parent.mkdir(parents=True, exist_ok=True)

    results: list[dict] = []
    print(f"\n=== Relative dip compare · S&P 500 · {args.start} → {args.end} ===\n", flush=True)

    for spec in variants:
        print(f"--- {spec.label} ---", flush=True)
        r, meta = _run_variant(
            spec,
            equity_dict=equity_dict,
            spy_df=spy_df,
            top_n=int(args.top_n),
            hold_days=int(args.hold_days),
            pct_drop=float(args.pct_drop_min),
            start=str(args.start),
            end=str(args.end),
            cap=cap,
        )
        daily_path = Path(f"{prefix}_{spec.key}_daily.csv")
        pd.DataFrame(
            {
                "date": r.index.strftime("%Y-%m-%d"),
                "daily_ret": r.values,
                "daily_pnl_usd": (r * cap).values,
                "equity_usd": (cap * (1.0 + r).cumprod()).values,
            }
        ).to_csv(daily_path, index=False)
        meta["daily_csv"] = str(daily_path)
        results.append(meta)
        print(
            f"  Return {meta['total_return_pct']:+.1f}%  CAGR {meta['cagr_pct']:+.1f}%  "
            f"Sharpe {meta['sharpe']:.2f}  maxDD {meta['max_drawdown_pct']:.1f}%  "
            f"β(SPY) {meta['beta_vs_spy']:.2f}  ρ(SPY) {meta['corr_vs_spy']:.2f}",
            flush=True,
        )
        for yr in ("2020", "2022"):
            yret = meta["yearly_return_pct_chained"].get(yr)
            if yret is not None:
                print(f"    {yr}: {yret:+.1f}%", flush=True)
        print(f"    → {daily_path}\n", flush=True)

    cmd = (
        f"cd {_REPO} && PYTHONUNBUFFERED=1 .venv/bin/python "
        f"RenTech/strategy_stack/run_relative_dip_compare.py "
        f"--start {args.start} --end {args.end} --capital {cap:.0f}"
    )
    summary = {
        "universe": "sp500",
        "n_tickers": len(equity_dict),
        "capital_usd": cap,
        "command": cmd,
        "variants": results,
    }
    json_path = Path(f"{prefix}_summary.json")
    json_path.write_text(json.dumps(summary, indent=2) + "\n", encoding="utf-8")

    md_lines = [
        f"# Relative dip compare ({args.start} → {args.end})",
        "",
        f"Universe: S&P 500 (n={len(equity_dict)}), top {args.top_n}, hold {args.hold_days}d, "
        f"pct_drop≥{args.pct_drop_min:.0%}, relative underperf≥{rel_min:.0%} (relative/hedged).",
        "",
        "| Variant | Return | CAGR | Sharpe | Max DD | β(SPY) | ρ(SPY) | 2020 | 2022 |",
        "|---------|--------|------|--------|--------|--------|--------|------|------|",
    ]
    for m in results:
        y20 = m["yearly_return_pct_chained"].get("2020", float("nan"))
        y22 = m["yearly_return_pct_chained"].get("2022", float("nan"))
        md_lines.append(
            f"| {m['variant']} | {m['total_return_pct']:+.1f}% | {m['cagr_pct']:+.1f}% | "
            f"{m['sharpe']:.2f} | {m['max_drawdown_pct']:.1f}% | {m['beta_vs_spy']:.2f} | "
            f"{m['corr_vs_spy']:.2f} | {y20:+.1f}% | {y22:+.1f}% |"
        )
    md_lines.extend(["", f"```bash\n{cmd}\n```", ""])
    md_path = Path(f"{prefix}_summary.md")
    md_path.write_text("\n".join(md_lines) + "\n", encoding="utf-8")

    print("=== Summary ===", flush=True)
    print(md_path.read_text(), flush=True)
    print(f"Wrote {json_path}", flush=True)


if __name__ == "__main__":
    main()
