#!/usr/bin/env python3
"""
Sweep **bear hedge** SPY regime / exit rules on SH (2016+).

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_ma_slope_inverse_spy_sweep.py --start 2016-01-04
"""

from __future__ import annotations

import argparse
import itertools
import json
import sys
from dataclasses import asdict
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.ma_slope_inverse_sleeve import MaSlopeInverseConfig, MaSlopeInverseSleeve
from RenTech.strategy_stack.main import _compute_daily_backtest_features
from RenTech.strategy_stack.run_ma_slope_inverse_spy_standard import _load_etf_dict, _metrics

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


def _grid(*, quick: bool) -> list[MaSlopeInverseConfig]:
    regimes = ["bear_dual_or_sma200", "bear_dual_slope", "below_sma200"] if not quick else [
        "bear_dual_or_sma200", "below_sma200",
    ]
    exits = [
        "spy_bull_both",
        "spy_regime_off",
        "spy_slope_or_sma200",
        "spy_fast_slope_positive",
    ] if not quick else ["spy_bull_both", "spy_slope_or_sma200"]
    inv_req = [False, True] if not quick else [False]
    tickers_opts = [("SH",), ("SH", "SDS", "SPXU")] if not quick else [("SH",)]
    configs: list[MaSlopeInverseConfig] = []
    for regime, exit_, req, tickers in itertools.product(regimes, exits, inv_req, tickers_opts):
        configs.append(
            MaSlopeInverseConfig(
                spy_regime=regime,
                spy_exit=exit_,
                require_inverse_momentum=req,
                tickers_preferred=tickers,
                selection="top_n" if len(tickers) == 1 else "all_active",
            )
        )
    # legacy baseline
    configs.append(
        MaSlopeInverseConfig(
            spy_regime="none",
            spy_exit="inverse_signal",
            require_inverse_momentum=True,
            tickers_preferred=("SH", "SDS", "SPXU"),
            selection="all_active",
        )
    )
    return configs


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--start", default="2016-01-04")
    ap.add_argument("--end", default="")
    ap.add_argument("--yahoo-period", default="max")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--quick", action="store_true")
    ap.add_argument("--out-prefix", type=Path, default=LOGS / "ma_slope_inverse_spy_sweep")
    args = ap.parse_args()

    etf_dict = _load_etf_dict(["SH", "SDS", "SPXU"], args.yahoo_period)
    spy_df = _compute_daily_backtest_features(
        DataLoader().fetch_daily("SPY", period=args.yahoo_period)
    )
    spy_df.index = pd.to_datetime(spy_df.index).tz_localize(None)
    spy_r = spy_df["ret"].astype(float)

    rows: list[dict] = []
    for i, cfg in enumerate(_grid(quick=args.quick)):
        tickers = {t: etf_dict[t] for t in cfg.tickers_preferred if t in etf_dict}
        if not tickers:
            tickers = etf_dict
        eng = MaSlopeInverseSleeve(config=cfg)
        r_full = eng.generate_returns(tickers, spy_df, verbose=False)
        r_full.index = pd.to_datetime(r_full.index).tz_localize(None)
        mask = r_full.index >= pd.Timestamp(args.start)
        if args.end.strip():
            mask &= r_full.index <= pd.Timestamp(args.end)
        r = r_full.loc[mask].fillna(0.0)
        m = _metrics(r, spy_r, float(args.capital))
        rows.append({"config_slug": cfg.slug(), **{f"cfg_{k}": v for k, v in asdict(cfg).items()}, **m})
        if (i + 1) % 10 == 0 or i + 1 == len(_grid(quick=args.quick)):
            print(f"  [{i+1}/{len(_grid(quick=args.quick))}] done", flush=True)

    df = pd.DataFrame(rows).sort_values("ret_2022_pct", ascending=False)
    slug = "quick" if args.quick else "full"
    prefix = args.out_prefix.expanduser().resolve()
    prefix.parent.mkdir(parents=True, exist_ok=True)
    csv_path = Path(f"{prefix}_{slug}.csv")
    df.to_csv(csv_path, index=False)

    top22 = df.nlargest(8, "ret_2022_pct")
    top_sh = df.nlargest(8, "sharpe")
    md = [
        f"# Bear hedge sweep ({slug})",
        "",
        f"CSV: `{csv_path}`",
        "",
        "## Top 2022 return",
        "",
        top22[["config_slug", "ret_2022_pct", "sharpe", "cagr_pct", "max_dd_pct", "corr_vs_spy"]].to_string(index=False),
        "",
        "## Top Sharpe (full sample)",
        "",
        top_sh[["config_slug", "sharpe", "ret_2022_pct", "cagr_pct", "max_dd_pct"]].to_string(index=False),
    ]
    md_path = Path(f"{prefix}_{slug}.md")
    md_path.write_text("\n".join(md) + "\n")
    print(f"\nWrote {csv_path}")
    print(f"Best 2022: {df.iloc[0]['config_slug']} -> {df.iloc[0]['ret_2022_pct']:.1f}%")


if __name__ == "__main__":
    main()
