#!/usr/bin/env python3
"""
Grid-search **MA / EMA slope momentum** on daily bars.

Tests whether steeper moving-average slope improves trend-follow returns on
SPY (default) or a comma-separated ticker list.

Stages (use ``--stage``):
  1 — MA type × period × slope lookback (coarse, long-only, pct slope)
  2 — Entry threshold × exit mode × price filter (refine around stage-1 winners)
  3 — Slope method × direction (long-short) × min hold
  4 — Dual-timeframe × exit (slope / ATR chandelier) × regime gate × vol sizing

Example::

    cd /Users/robzingale/trading_bot
    PYTHONUNBUFFERED=1 .venv/bin/python RenTech/strategy_stack/run_ma_slope_sweep.py \\
        --start 2016-01-04 --stage 4 --quick

Full stage-1::

    PYTHONUNBUFFERED=1 .venv/bin/python RenTech/strategy_stack/run_ma_slope_sweep.py \\
        --start 2016-01-04 --stage 1

Single config (export daily CSV)::

    PYTHONUNBUFFERED=1 .venv/bin/python RenTech/strategy_stack/run_ma_slope_sweep.py \\
        --start 2016-01-04 --single \\
        --ma-type ema --ma-period 50 --slope-lookback 10 \\
        --entry-slope-min 0.01 --export-daily
"""

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_engine import (
    MaSlopeConfig,
    MaSlopeEngine,
    metrics_from_returns,
)
from RenTech.strategy_stack.main import _compute_daily_backtest_features

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


def _load_vix_series(start: str) -> pd.Series:
    import yfinance as yf

    warmup = (pd.Timestamp(start) - pd.Timedelta(days=400)).strftime("%Y-%m-%d")
    raw = yf.download("^VIX", start=warmup, progress=False, auto_adjust=True)
    if raw.empty:
        raise RuntimeError("Failed to download ^VIX from Yahoo Finance")
    if isinstance(raw.columns, pd.MultiIndex):
        raw = raw.droplevel(1, axis=1)
    col = "Close" if "Close" in raw.columns else "close"
    vix = raw[col].astype(np.float64)
    vix.index = pd.to_datetime(vix.index).tz_localize(None)
    return vix


def _attach_vix(df: pd.DataFrame, vix: pd.Series) -> pd.DataFrame:
    out = df.copy()
    out["vix"] = vix.reindex(out.index).ffill()
    return out


def _load_ticker(daily: pd.DataFrame, start: str, end: str) -> pd.DataFrame:
    df = _compute_daily_backtest_features(daily)
    df.index = pd.to_datetime(df.index).tz_localize(None)
    mask = df.index >= pd.Timestamp(start)
    if end.strip():
        mask &= df.index <= pd.Timestamp(end)
    return df.loc[mask].copy()


def _period_return(r: pd.Series, start: str, end: str) -> float:
    sub = r.loc[start:end]
    if len(sub) < 2:
        return float("nan")
    return float((1.0 + sub).prod() - 1.0) * 100.0


def _run_config(
    df: pd.DataFrame,
    cfg: MaSlopeConfig,
    *,
    capital: float,
    benchmark_ret: pd.Series,
) -> dict:
    eng = MaSlopeEngine(config=cfg)
    frame = eng.transform(df)
    strat_r = frame["position"].astype(np.float64) * df["ret"].astype(np.float64)
    strat_r = strat_r.fillna(0.0)
    m = metrics_from_returns(strat_r, capital=capital, position=frame["position"])
    aligned = pd.DataFrame({"s": strat_r, "b": benchmark_ret}).dropna()
    rho = float(aligned.corr().iloc[0, 1]) if len(aligned) > 2 else float("nan")
    row = {
        "config_slug": cfg.slug(),
        **{f"cfg_{k}": v for k, v in asdict(cfg).items()},
        **m,
        "beta_spy": rho * (m.get("vol_ann_pct", float("nan")) / (benchmark_ret.std(ddof=1) * np.sqrt(252) * 100 + 1e-12)),
        "ret_2022_pct": _period_return(strat_r, "2022-01-01", "2022-12-31"),
        "ret_2020_pct": _period_return(strat_r, "2020-01-01", "2020-12-31"),
        "ret_2018_pct": _period_return(strat_r, "2018-01-01", "2018-12-31"),
    }
    return row


def _stage1_grid(*, quick: bool) -> list[MaSlopeConfig]:
    ma_types = ["sma", "ema"]
    periods = [20, 50, 100, 200] if quick else [10, 20, 50, 100, 200]
    lookbacks = [5, 10, 20] if quick else [5, 10, 20, 40]
    # Entry threshold scales with lookback: ~1% per 10 bars is a reasonable prior.
    entry_map = {5: 0.005, 10: 0.01, 20: 0.02, 40: 0.03}
    configs: list[MaSlopeConfig] = []
    for ma_type, period, lb in itertools.product(ma_types, periods, lookbacks):
        entry = entry_map.get(lb, 0.01)
        configs.append(
            MaSlopeConfig(
                ma_type=ma_type,
                ma_period=period,
                slope_lookback=lb,
                slope_method="pct",
                entry_slope_min=entry,
                price_above_ma=True,
                direction="long_only",
                exit_mode="slope_flip",
            )
        )
    return configs


def _stage2_grid(*, quick: bool) -> list[MaSlopeConfig]:
    """Refinement around stage-1 winners (fast MAs: 10–20 day)."""
    bases = [
        ("ema", 10, 10),
        ("ema", 10, 5),
        ("sma", 10, 10),
        ("ema", 20, 5),
    ] if quick else [
        ("ema", 10, 10),
        ("ema", 10, 5),
        ("ema", 10, 20),
        ("sma", 10, 10),
        ("sma", 10, 5),
        ("ema", 20, 5),
        ("ema", 20, 10),
        ("sma", 20, 5),
    ]
    entries = [0.0, 0.005, 0.01, 0.015, 0.02, 0.03] if not quick else [0.005, 0.01, 0.02]
    exits: list[tuple[str, float | None]] = [
        ("slope_flip", None),
        ("slope_half", None),
        ("price_cross", None),
        ("slope_or_price", None),
    ]
    price_filters = [True, False]
    configs: list[MaSlopeConfig] = []
    for (ma_type, period, lb), entry, (exit_mode, exit_max), pf in itertools.product(
        bases, entries, exits, price_filters
    ):
        configs.append(
            MaSlopeConfig(
                ma_type=ma_type,
                ma_period=period,
                slope_lookback=lb,
                slope_method="pct",
                entry_slope_min=entry,
                exit_slope_max=exit_max,
                price_above_ma=pf,
                direction="long_only",
                exit_mode=exit_mode,
            )
        )
    return configs


def _stage3_grid(*, quick: bool) -> list[MaSlopeConfig]:
    methods = ["pct", "annualized_pct"] if quick else ["pct", "annualized_pct", "regression"]
    directions = ["long_only", "long_short"]
    holds = [0, 5] if quick else [0, 3, 5, 10]
    bases = [
        MaSlopeConfig(ma_type="ema", ma_period=10, slope_lookback=10, entry_slope_min=0.01),
        MaSlopeConfig(ma_type="ema", ma_period=10, slope_lookback=5, entry_slope_min=0.005),
        MaSlopeConfig(ma_type="sma", ma_period=10, slope_lookback=10, entry_slope_min=0.01),
    ]
    configs: list[MaSlopeConfig] = []
    for base, method, direction, hold in itertools.product(bases, methods, directions, holds):
        cfg = MaSlopeConfig(
            ma_type=base.ma_type,
            ma_period=base.ma_period,
            slope_lookback=base.slope_lookback,
            slope_method=method,
            entry_slope_min=base.entry_slope_min,
            price_above_ma=True,
            direction=direction,
            exit_mode="slope_or_price",
            min_hold_days=hold,
        )
        configs.append(cfg)
    return configs


def _stage4_base_config() -> MaSlopeConfig:
    """Stage-1/2 SPY winner: EMA(10), slope LB 10, zero entry threshold."""
    return MaSlopeConfig(
        ma_type="ema",
        ma_period=10,
        slope_lookback=10,
        slope_method="pct",
        entry_slope_min=0.0,
        price_above_ma=True,
        direction="long_only",
        exit_mode="slope_flip",
        signal_mode="dual_timeframe",
    )


def _stage4_grid(*, quick: bool) -> list[MaSlopeConfig]:
    base = _stage4_base_config()
    slow_periods = [50] if quick else [50, 100]
    slow_lbs = [10] if quick else [5, 10, 20]
    exit_variants: list[tuple[str, float | None]] = [
        ("slope_flip", None),
        ("slope_or_price", None),
    ]
    atr_mults = [2.5] if quick else [2.0, 2.5, 3.0]
    regime_gates = ["none", "sma200"] if quick else ["none", "sma200", "vix_lt"]
    sizing_variants: list[tuple[str, float | None]] = [("binary", None)]
    if quick:
        sizing_variants.append(("vol_scale", 0.15))
    else:
        sizing_variants.extend([("vol_scale", v) for v in (0.10, 0.15, 0.20)])

    configs: list[MaSlopeConfig] = []
    seen: set[str] = set()

    def add(cfg: MaSlopeConfig) -> None:
        slug = cfg.slug()
        if slug not in seen:
            seen.add(slug)
            configs.append(cfg)

    for slow_p, slow_lb, (exit_mode, _), regime, (sz_mode, tgt_vol) in itertools.product(
        slow_periods, slow_lbs, exit_variants, regime_gates, sizing_variants
    ):
        add(
            MaSlopeConfig(
                ma_type=base.ma_type,
                ma_period=base.ma_period,
                slope_lookback=base.slope_lookback,
                slope_method=base.slope_method,
                entry_slope_min=base.entry_slope_min,
                price_above_ma=base.price_above_ma,
                direction=base.direction,
                exit_mode=exit_mode,
                signal_mode="dual_timeframe",
                slow_ma_period=slow_p,
                slow_slope_lookback=slow_lb,
                regime_gate=regime,
                sizing_mode=sz_mode,
                target_vol=tgt_vol if tgt_vol is not None else base.target_vol,
            )
        )

    for slow_p, slow_lb, atr_mult, regime, (sz_mode, tgt_vol) in itertools.product(
        slow_periods, slow_lbs, atr_mults, regime_gates, sizing_variants
    ):
        add(
            MaSlopeConfig(
                ma_type=base.ma_type,
                ma_period=base.ma_period,
                slope_lookback=base.slope_lookback,
                slope_method=base.slope_method,
                entry_slope_min=base.entry_slope_min,
                price_above_ma=base.price_above_ma,
                direction=base.direction,
                exit_mode="atr_chandelier",
                signal_mode="dual_timeframe",
                slow_ma_period=slow_p,
                slow_slope_lookback=slow_lb,
                regime_gate=regime,
                atr_multiplier=atr_mult,
                sizing_mode=sz_mode,
                target_vol=tgt_vol if tgt_vol is not None else base.target_vol,
            )
        )

    return configs


def _grid_for_stage(stage: int, *, quick: bool) -> list[MaSlopeConfig]:
    if stage == 1:
        return _stage1_grid(quick=quick)
    if stage == 2:
        return _stage2_grid(quick=quick)
    if stage == 3:
        return _stage3_grid(quick=quick)
    if stage == 4:
        return _stage4_grid(quick=quick)
    raise ValueError(f"stage must be 1–4; got {stage}")


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="")
    ap.add_argument("--ticker", default="SPY", help="Symbol or comma-separated list (first is benchmark)")
    ap.add_argument("--yahoo-period", default="max")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--stage", type=int, default=1, choices=(1, 2, 3, 4))
    ap.add_argument("--quick", action="store_true", help="Smaller grid for smoke tests")
    ap.add_argument("--single", action="store_true", help="Run one config from CLI flags")
    ap.add_argument("--ma-type", default="ema", choices=("sma", "ema"))
    ap.add_argument("--ma-period", type=int, default=50)
    ap.add_argument("--slope-lookback", type=int, default=10)
    ap.add_argument("--slope-method", default="pct", choices=("pct", "annualized_pct", "regression"))
    ap.add_argument("--entry-slope-min", type=float, default=0.01)
    ap.add_argument("--exit-slope-max", type=float, default=None)
    ap.add_argument("--no-price-filter", action="store_true")
    ap.add_argument("--direction", default="long_only", choices=("long_only", "long_short"))
    ap.add_argument(
        "--exit-mode",
        default="slope_flip",
        choices=("slope_flip", "slope_half", "price_cross", "slope_or_price", "atr_chandelier"),
    )
    ap.add_argument("--min-hold-days", type=int, default=0)
    ap.add_argument("--signal-mode", default="single", choices=("single", "dual_timeframe"))
    ap.add_argument("--slow-ma-period", type=int, default=50)
    ap.add_argument("--slow-slope-lookback", type=int, default=None)
    ap.add_argument("--regime-gate", default="none", choices=("none", "sma200", "vix_lt"))
    ap.add_argument("--vix-cap", type=float, default=25.0)
    ap.add_argument("--sizing-mode", default="binary", choices=("binary", "vol_scale"))
    ap.add_argument("--target-vol", type=float, default=0.15)
    ap.add_argument("--atr-multiplier", type=float, default=2.5)
    ap.add_argument("--export-daily", action="store_true", help="With --single, write daily CSV")
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT)
    ap.add_argument("--top-n", type=int, default=25, help="Rows in markdown summary")
    args = ap.parse_args()

    tickers = [t.strip().upper() for t in args.ticker.split(",") if t.strip()]
    if not tickers:
        raise SystemExit("No tickers specified")

    loader = DataLoader()
    panels: dict[str, pd.DataFrame] = {}
    for t in tickers:
        raw = loader.fetch_daily(t, period=args.yahoo_period)
        panels[t] = _load_ticker(raw, args.start, args.end)

    if (not args.single and args.stage == 4) or (args.single and args.regime_gate == "vix_lt"):
        vix_series = _load_vix_series(args.start)
        for t in panels:
            panels[t] = _attach_vix(panels[t], vix_series)

    benchmark = tickers[0]
    bench_ret = panels[benchmark]["ret"].astype(np.float64)

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

    if args.single:
        cfg = MaSlopeConfig(
            ma_type=args.ma_type,
            ma_period=args.ma_period,
            slope_lookback=args.slope_lookback,
            slope_method=args.slope_method,
            entry_slope_min=args.entry_slope_min,
            exit_slope_max=args.exit_slope_max,
            price_above_ma=not args.no_price_filter,
            direction=args.direction,
            exit_mode=args.exit_mode,
            min_hold_days=args.min_hold_days,
            signal_mode=args.signal_mode,
            slow_ma_period=args.slow_ma_period,
            slow_slope_lookback=args.slow_slope_lookback,
            regime_gate=args.regime_gate,
            vix_cap=args.vix_cap,
            sizing_mode=args.sizing_mode,
            target_vol=args.target_vol,
            atr_multiplier=args.atr_multiplier,
        )
        grid = [cfg]
    else:
        grid = _grid_for_stage(args.stage, quick=args.quick)

    rows: list[dict] = []
    for i, cfg in enumerate(grid):
        for t, df in panels.items():
            row = _run_config(df, cfg, capital=args.capital, benchmark_ret=bench_ret)
            row["ticker"] = t
            rows.append(row)
        if (i + 1) % 20 == 0 or i + 1 == len(grid):
            print(f"  [{i + 1}/{len(grid)}] configs done", flush=True)

    result = pd.DataFrame(rows)
    result = result.sort_values(["ticker", "sharpe"], ascending=[True, False])

    slug = f"stage{args.stage}" + ("_quick" if args.quick else "")
    if args.single:
        slug = cfg.slug()
    csv_path = Path(f"{prefix}_{slug}.csv")
    result.to_csv(csv_path, index=False)

    md_path = Path(f"{prefix}_{slug}.md")
    top = result.head(int(args.top_n))
    cols = [
        "ticker",
        "config_slug",
        "sharpe",
        "cagr_pct",
        "max_dd_pct",
        "total_return_pct",
        "invested_frac",
        "ret_2022_pct",
    ]
    header = "| " + " | ".join(cols) + " |"
    sep = "| " + " | ".join("---" for _ in cols) + " |"
    body_lines = []
    for _, row in top.iterrows():
        body_lines.append(
            "| "
            + " | ".join(
                f"{row[c]:.2f}" if isinstance(row[c], (float, np.floating)) else str(row[c])
                for c in cols
            )
            + " |"
        )
    lines = [
        f"# MA slope sweep — stage {args.stage}" + (" (quick)" if args.quick else ""),
        "",
        f"- Window: `{args.start}` → `{args.end or 'latest'}`",
        f"- Tickers: `{', '.join(tickers)}`",
        f"- Configs: **{len(grid)}** × **{len(tickers)}** tickers = **{len(result)}** rows",
        f"- CSV: `{csv_path}`",
        "",
        "## Top by Sharpe",
        "",
        header,
        sep,
        *body_lines,
    ]
    md_path.write_text("\n".join(lines) + "\n")

    meta = {
        "start": args.start,
        "end": args.end or None,
        "tickers": tickers,
        "stage": args.stage,
        "quick": args.quick,
        "n_configs": len(grid),
        "csv": str(csv_path),
        "markdown": str(md_path),
    }
    meta_path = Path(f"{prefix}_{slug}_meta.json")
    meta_path.write_text(json.dumps(meta, indent=2) + "\n")

    print(f"\nWrote {csv_path}")
    print(f"Wrote {md_path}")
    print(f"Wrote {meta_path}")
    if len(result):
        best = result.iloc[0]
        print(
            f"\nBest ({best['ticker']}): Sharpe {best['sharpe']:.2f} · "
            f"CAGR {best['cagr_pct']:.1f}% · max DD {best['max_dd_pct']:.1f}% · "
            f"{best['config_slug']}"
        )

    if args.single and args.export_daily:
        eng = MaSlopeEngine(config=grid[0])
        frame = eng.transform(panels[benchmark])
        daily_path = Path(f"{prefix}_{slug}_daily.csv")
        out = pd.DataFrame(
            {
                "date": frame.index.strftime("%Y-%m-%d"),
                "close": panels[benchmark]["close"].values,
                "ret": panels[benchmark]["ret"].values,
                "ma": frame["ma"].values,
                "ma_slope": frame["ma_slope"].values,
                "position": frame["position"].values,
                "daily_ret": (
                    frame["position"].astype(np.float64) * panels[benchmark]["ret"].astype(np.float64)
                ).values,
            }
        )
        if "ma_slow" in frame.columns:
            out["ma_slow"] = frame["ma_slow"].values
            out["ma_slow_slope"] = frame["ma_slow_slope"].values
        if "position_weight" in frame.columns:
            out["position_weight"] = frame["position_weight"].values
        out["equity_usd"] = float(args.capital) * (1.0 + out["daily_ret"]).cumprod()
        out.to_csv(daily_path, index=False)
        print(f"Wrote {daily_path}")


if __name__ == "__main__":
    main()
