#!/usr/bin/env python3
"""
Grid **buy-the-dip** sleeves across universes (S&P 500, S&P 100, SPDR sectors).

Writes ranked summary CSV + Markdown under ``RenTech/data/logs/``.

Example::

    cd /Users/robzingale/trading_bot
    PYTHONUNBUFFERED=1 .venv/bin/python RenTech/strategy_stack/experiment_buy_the_dip.py \\
        --universe sp500 --start 2016-01-04 --yahoo-period max
"""

from __future__ import annotations

import argparse
import json
import sys
from dataclasses import dataclass, fields
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))

from RenTech.strategy_stack.data_loader import DataLoader
from RenTech.strategy_stack.main import (
    _compute_daily_backtest_features,
    _load_aqr_equity_dict,
    _load_sector_etf_dict,
)
from RenTech.strategy_stack.multi_strategy_manager import BuyTheDipSleeve

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


@dataclass(frozen=True)
class DipVariant:
    label: str
    universe: str  # sp500 | sp100 | sector
    top_n: int = 10
    rsi_max: float = 20.0
    hold_trading_days: int = 5
    signal_mode: str = "rsi"
    pct_drop_min: float = 0.03
    min_atr_norm_pct: float = 3.0
    rank_by: str = "atr_norm"
    weighting: str = "equal"
    only_when_spy_bull: bool = False
    dip_in_uptrend: bool = True

    def to_engine(self) -> BuyTheDipSleeve:
        return BuyTheDipSleeve(
            rsi_max=float(self.rsi_max),
            hold_trading_days=int(self.hold_trading_days),
            rank_by=str(self.rank_by),
            weighting=str(self.weighting),
            dip_in_uptrend=bool(self.dip_in_uptrend),
            only_when_spy_bull=bool(self.only_when_spy_bull),
            signal_mode=str(self.signal_mode),
            pct_drop_min=float(self.pct_drop_min),
            min_atr_norm_pct=float(self.min_atr_norm_pct),
        )


def _metrics(r: pd.Series) -> dict[str, Any]:
    r = r.fillna(0.0).astype(np.float64)
    n = len(r)
    if n < 2:
        return {"n": n}
    eq = (1.0 + r).cumprod()
    years = n / 252.0
    end = float(eq.iloc[-1])
    total_ret = end - 1.0
    cagr = end ** (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")
    max_dd = float((eq / eq.cummax() - 1.0).min())
    active = float((r.abs() > 1e-10).mean())
    return {
        "n": int(n),
        "total_return_pct": round(100.0 * total_ret, 2),
        "cagr_pct": round(100.0 * cagr, 2),
        "sharpe": round(sharpe, 3),
        "vol_ann_pct": round(100.0 * sd * np.sqrt(252.0), 2),
        "max_dd_pct": round(100.0 * max_dd, 2),
        "active_days_pct": round(100.0 * active, 1),
    }


def _corr_vs(path: Path, r: pd.Series, col: str = "daily_ret") -> float | None:
    if not path.is_file():
        return None
    other = pd.read_csv(path, parse_dates=["date"]).set_index("date")[col].astype(float)
    other.index = pd.to_datetime(other.index).tz_localize(None)
    both = pd.DataFrame({"a": r, "b": other}).dropna()
    if len(both) < 30:
        return None
    return float(both["a"].corr(both["b"]))


def _build_variants(universe: str) -> list[DipVariant]:
    u = universe.lower()
    if u == "sector":
        top_defaults = (1, 3, 5)
    elif u in ("russell3000", "r3k"):
        top_defaults = (5, 10, 20)
    else:
        top_defaults = (5, 10, 20)

    variants: list[DipVariant] = []
    for top_n in top_defaults:
        variants.append(
            DipVariant(
                label=f"{u}_rsi20_hold5_top{top_n}",
                universe=u,
                top_n=top_n,
            )
        )
    variants.extend(
        [
            DipVariant(
                label=f"{u}_rsi20_hold10_top10",
                universe=u,
                top_n=10 if u != "sector" else 3,
                hold_trading_days=10,
            ),
            DipVariant(
                label=f"{u}_rsi15_hold5_top10",
                universe=u,
                top_n=10 if u != "sector" else 3,
                rsi_max=15.0,
            ),
            DipVariant(
                label=f"{u}_rsi25_hold5_top10",
                universe=u,
                top_n=10 if u != "sector" else 3,
                rsi_max=25.0,
            ),
            DipVariant(
                label=f"{u}_rsi20_spy_bull_top10",
                universe=u,
                top_n=10 if u != "sector" else 3,
                only_when_spy_bull=True,
            ),
            DipVariant(
                label=f"{u}_rsi20_invvol_top10",
                universe=u,
                top_n=10 if u != "sector" else 3,
                weighting="inv_vol",
            ),
            DipVariant(
                label=f"{u}_rsi20_rank_rsi_top10",
                universe=u,
                top_n=10 if u != "sector" else 3,
                rank_by="rsi",
            ),
            DipVariant(
                label=f"{u}_pct3_hold5_top10",
                universe=u,
                top_n=10 if u != "sector" else 3,
                signal_mode="pct_drop",
                pct_drop_min=0.03,
            ),
            DipVariant(
                label=f"{u}_pct3_hold10_top10",
                universe=u,
                top_n=10 if u != "sector" else 3,
                signal_mode="pct_drop",
                hold_trading_days=10,
            ),
            DipVariant(
                label=f"{u}_pct5_hold5_top10",
                universe=u,
                top_n=10 if u != "sector" else 3,
                signal_mode="pct_drop",
                pct_drop_min=0.05,
            ),
        ]
    )
    if u != "sector":
        variants.append(
            DipVariant(
                label=f"{u}_rsi20_no_trend_top10",
                universe=u,
                dip_in_uptrend=False,
            )
        )
    return variants


def _load_universe(
    universe: str,
    yahoo_period: str,
    max_tickers: int,
) -> tuple[dict[str, pd.DataFrame], str]:
    u = universe.lower()
    if u == "sector":
        return _load_sector_etf_dict(yahoo_period), "sector_11_spdr"
    if u in ("sp500", "sp100"):
        eq = _load_aqr_equity_dict(
            yahoo_period,
            universe=u,
            max_tickers=max_tickers,
        )
        return eq, f"{u}_n{len(eq)}"
    if u in ("russell3000", "r3k"):
        from RenTech.strategy_stack.equity_universe_loaders import load_equity_panel_dict

        eq = load_equity_panel_dict("russell3000", yahoo_period, max_tickers=max_tickers)
        return eq, f"russell3000_n{len(eq)}"
    raise ValueError(f"universe must be sp500, sp100, russell3000, or sector; got {universe!r}")


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
    ap.add_argument(
        "--universe",
        default="sp500",
        choices=("sp500", "sp100", "russell3000", "sector", "all"),
    )
    ap.add_argument("--start", default="2016-01-04")
    ap.add_argument("--end", default="")
    ap.add_argument("--yahoo-period", default="10y")
    ap.add_argument("--max-tickers", type=int, default=0, help="0 = all names (S&P only)")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument(
        "--out-stem",
        type=Path,
        default=LOGS / "buy_the_dip_experiment",
    )
    ap.add_argument("--quiet", action="store_true")
    args = ap.parse_args()

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

    universes = (
        ["sp500", "sp100", "russell3000", "sector"]
        if args.universe == "all"
        else [args.universe]
    )
    rows: list[dict[str, Any]] = []
    cache: dict[str, dict[str, pd.DataFrame]] = {}

    for u in universes:
        if u not in cache:
            cache[u], uni_label = _load_universe(u, args.yahoo_period, args.max_tickers)
            print(f"Loaded {uni_label}: {len(cache[u])} tickers", flush=True)
        equity_dict = cache[u]
        if len(equity_dict) < 5:
            raise SystemExit(f"Too few tickers for universe {u}")

        for var in _build_variants(u):
            eng = var.to_engine()
            daily = eng.generate_returns(
                equity_dict,
                top_n=var.top_n,
                spy_df=spy_df,
                verbose=not args.quiet,
            )
            daily.index = pd.to_datetime(daily.index).tz_localize(None)
            mask = daily.index >= pd.Timestamp(args.start)
            if args.end.strip():
                mask &= daily.index <= pd.Timestamp(args.end)
            r = daily.loc[mask].fillna(0.0).astype(np.float64)

            m = _metrics(r)
            row = {
                "label": var.label,
                "universe": u,
                "n_tickers": len(equity_dict),
                **{f.name: getattr(var, f.name) for f in fields(var) if f.name not in ("label", "universe")},
                **m,
            }
            if u == "sector":
                row["corr_vs_sector_mom"] = _corr_vs(
                    LOGS / "sector_momentum_standard_daily.csv", r
                )
            if u == "sp500":
                row["corr_vs_sector_dip"] = _corr_vs(
                    LOGS / "sector_dip_standard_daily.csv", r
                )
            rows.append(row)
            print(
                f"  {var.label:32s}  ret={m.get('total_return_pct', float('nan')):7.2f}%  "
                f"Sharpe={m.get('sharpe', float('nan')):5.2f}  DD={m.get('max_dd_pct', float('nan')):6.2f}%",
                flush=True,
            )

    df = pd.DataFrame(rows)
    if df.empty:
        raise SystemExit("No results")
    df = df.sort_values(["universe", "sharpe"], ascending=[True, False], na_position="last")

    stem = args.out_stem.expanduser().resolve()
    if len(universes) == 1:
        stem = Path(f"{stem}_{universes[0]}_{args.start.replace('-', '')}")
    else:
        stem = Path(f"{stem}_all_{args.start.replace('-', '')}")

    csv_path = Path(f"{stem}.csv")
    md_path = Path(f"{stem}.md")
    json_path = Path(f"{stem}.json")
    df.to_csv(csv_path, index=False)

    best = df.groupby("universe", sort=False).first()
    md_lines = [
        f"# Buy-the-dip experiment ({args.start}" + (f" → {args.end}" if args.end else "") + ")",
        "",
        f"Yahoo period: `{args.yahoo_period}` · Capital notional: ${args.capital:,.0f} (return stream is unit-scaled; compare Sharpe/DD).",
        "",
        "## Top variant per universe (by Sharpe)",
        "",
        "| Universe | Label | Return % | CAGR % | Sharpe | Max DD % | Active % |",
        "|----------|-------|----------|--------|--------|----------|----------|",
    ]
    for u, row in best.iterrows():
        md_lines.append(
            f"| {u} | {row['label']} | {row['total_return_pct']:.1f} | {row['cagr_pct']:.1f} | "
            f"{row['sharpe']:.2f} | {row['max_dd_pct']:.1f} | {row.get('active_days_pct', float('nan')):.1f} |"
        )
    md_lines.extend(["", "## Full grid (sorted by Sharpe within universe)", ""])
    for u in universes:
        sub = df[df["universe"] == u].sort_values("sharpe", ascending=False)
        md_lines.append(f"### {u}")
        md_lines.append("")
        md_lines.append(
            "| Label | Return % | Sharpe | Max DD % | Hold | Top N | Signal |"
        )
        md_lines.append("|-------|----------|--------|----------|------|-------|--------|")
        for _, row in sub.iterrows():
            sig = row["signal_mode"]
            if sig == "rsi":
                sig = f"rsi<{row['rsi_max']:.0f}"
            else:
                sig = f"drop>={row['pct_drop_min']:.0%}"
            md_lines.append(
                f"| {row['label']} | {row['total_return_pct']:.1f} | {row['sharpe']:.2f} | "
                f"{row['max_dd_pct']:.1f} | {int(row['hold_trading_days'])} | {int(row['top_n'])} | {sig} |"
            )
        md_lines.append("")

    md_lines.extend(
        [
            "## Reproduce",
            "",
            "```bash",
            f"cd {_REPO} && PYTHONUNBUFFERED=1 .venv/bin/python "
            f"RenTech/strategy_stack/experiment_buy_the_dip.py "
            f"--universe {args.universe} --start {args.start} --yahoo-period {args.yahoo_period}",
            "```",
            "",
            f"CSV: `{csv_path}`",
        ]
    )
    md_path.write_text("\n".join(md_lines) + "\n")
    json_path.write_text(
        json.dumps(
            {
                "start": args.start,
                "end": args.end or None,
                "yahoo_period": args.yahoo_period,
                "universes": universes,
                "n_variants": len(df),
                "csv": str(csv_path),
            },
            indent=2,
        )
        + "\n"
    )

    print(f"\nWrote {csv_path}", flush=True)
    print(f"Wrote {md_path}", flush=True)
    print("\n--- Top by universe ---", flush=True)
    print(best[["label", "total_return_pct", "sharpe", "max_dd_pct"]].to_string(), flush=True)


if __name__ == "__main__":
    main()
