#!/usr/bin/env python3
"""
Backtest QuantifiedStrategies-style **top-tier** ETF sleeves (2016+ window).

Families (one rule set each, plain-English QS conventions):
  1. seasonality_turnaround_tuesday   — long SPY Tue if Mon return < 0
  2. seasonality_turn_of_month        — long last 1 + first 3 sessions of month
  3. seasonality_opex_week            — long Mon–Fri of week with 3rd Friday
  4. overnight_close_to_open          — long SPY every close → next open
  5. overnight_5d_low                 — close at 5-day low → next open
  6. overnight_3down                  — after 3 down days → next open
  7. mr_rsi2_sma200                   — Connors RSI(2)<10, >SMA200; exit RSI>70 or 5d
  8. mr_ibs_sma200                    — IBS<0.2, >SMA200; exit IBS>0.8 or 3d
  9. dual_momentum_spy_tlt            — Antonacci 12m abs+rel mom (SPY vs TLT)
 10. rotation_spy_tlt_gld             — monthly top-1 12m momentum

Example::

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

from __future__ import annotations

import argparse
import json
import sys
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 _standardize_ohlcv

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

try:
    import yfinance as yf
except ImportError:
    yf = None  # type: ignore


def _fetch_ohlc(ticker: str, period: str) -> pd.DataFrame:
    """Adjusted OHLC (auto_adjust=True) so open/close are on one scale for overnight math."""
    if yf is None:
        raise ImportError("yfinance required")
    raw = yf.Ticker(ticker).history(period=period, interval="1d", auto_adjust=True)
    out = _standardize_ohlcv(raw, prefer_adjusted=False)
    out.index = pd.to_datetime(out.index).tz_localize(None)
    return out.sort_index()


def _rsi(close: pd.Series, n: int) -> pd.Series:
    d = close.diff()
    up = d.clip(lower=0.0)
    dn = (-d).clip(lower=0.0)
    avg_up = up.ewm(alpha=1.0 / n, min_periods=n, adjust=False).mean()
    avg_dn = dn.ewm(alpha=1.0 / n, min_periods=n, adjust=False).mean()
    rs = avg_up / avg_dn.replace(0, np.nan)
    return 100.0 - (100.0 / (1.0 + rs))


def _ibs(df: pd.DataFrame) -> pd.Series:
    rng = (df["high"] - df["low"]).replace(0, np.nan)
    return (df["close"] - df["low"]) / rng


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": n, "total_return_pct": 0.0, "cagr_pct": 0.0, "sharpe": 0.0, "max_dd_pct": 0.0}
    eq = capital * (1.0 + r).cumprod()
    years = n / 252.0
    end_eq = float(eq.iloc[-1])
    total_ret = end_eq / capital - 1.0
    cagr = (end_eq / capital) ** (1.0 / years) - 1.0 if years > 0 else float("nan")
    dd = eq / eq.cummax() - 1.0
    sd = float(r.std(ddof=1))
    sharpe = float(r.mean() / sd * np.sqrt(252.0)) if sd > 1e-12 else 0.0
    return {
        "n_days": n,
        "total_return_pct": round(total_ret * 100.0, 2),
        "cagr_pct": round(cagr * 100.0, 2),
        "sharpe": round(sharpe, 3),
        "max_dd_pct": round(float(dd.min()) * 100.0, 2),
        "vol_ann_pct": round(sd * np.sqrt(252.0) * 100.0, 2),
        "ending_equity_usd": round(end_eq, 2),
        "pct_days_invested": round(float((r != 0.0).mean()) * 100.0, 1),
    }


def _align(df: pd.DataFrame, start: pd.Timestamp, end: pd.Timestamp | None) -> pd.DataFrame:
    out = df.sort_index()
    out.index = pd.to_datetime(out.index).tz_localize(None)
    mask = out.index >= start
    if end is not None:
        mask &= out.index <= end
    return out.loc[mask].copy()


def _spy_close_to_close(spy: pd.DataFrame) -> pd.Series:
    return spy["close"].pct_change().fillna(0.0)


def seasonality_turnaround_tuesday(spy: pd.DataFrame) -> pd.Series:
    """Long SPY on Tuesday when Monday's close-to-close return is negative."""
    c2c = _spy_close_to_close(spy)
    mon_neg = (spy.index.dayofweek == 0) & (c2c.shift(0) < 0)  # Monday signal on Mon
    # Position on Tuesday: use Monday's down day as signal
    signal = pd.Series(False, index=spy.index)
    for i in range(1, len(spy)):
        if spy.index[i].dayofweek == 1 and spy.index[i - 1].dayofweek == 0:
            if c2c.iloc[i - 1] < 0:
                signal.iloc[i] = True
    return c2c.where(signal, 0.0)


def seasonality_turn_of_month(spy: pd.DataFrame) -> pd.Series:
    """Long last session of month + first 3 sessions of next month (QS turn-of-month)."""
    idx = spy.index
    in_window = pd.Series(False, index=idx)
    months = pd.Series(idx, index=idx).dt.to_period("M")
    for per, dates in months.groupby(months):
        dlist = list(dates.index)
        if not dlist:
            continue
        picks = {dlist[-1]}
        # first 3 of this month
        for d in dlist[:3]:
            picks.add(d)
        for d in picks:
            in_window.loc[d] = True
    return _spy_close_to_close(spy).where(in_window, 0.0)


def _third_friday(year: int, month: int) -> pd.Timestamp:
    d = pd.Timestamp(year=year, month=month, day=1)
    fridays = pd.date_range(d, d + pd.offsets.MonthEnd(0), freq="W-FRI")
    return fridays[2] if len(fridays) >= 3 else fridays[-1]


def seasonality_opex_week(spy: pd.DataFrame) -> pd.Series:
    """Long Mon–Fri of calendar week containing the 3rd Friday (OPEX week)."""
    idx = spy.index
    opex_days: set[pd.Timestamp] = set()
    for per in idx.to_period("M").unique():
        tf = _third_friday(per.year, per.month)
        week_start = tf - pd.Timedelta(days=tf.dayofweek)
        for k in range(5):
            opex_days.add((week_start + pd.Timedelta(days=k)).normalize())
    in_week = pd.Series([d.normalize() in opex_days for d in idx], index=idx)
    return _spy_close_to_close(spy).where(in_week, 0.0)


def overnight_close_to_open(spy: pd.DataFrame) -> pd.Series:
    """Buy every close, sell next open; PnL booked on next session date."""
    oc = spy["open"] / spy["close"].shift(1) - 1.0
    return oc.fillna(0.0)


def _overnight_one_night(spy: pd.DataFrame, entry_mask: pd.Series) -> pd.Series:
    """Enter at signal-day close, exit next open; return on exit day."""
    oc = spy["open"] / spy["close"].shift(1) - 1.0
    out = pd.Series(0.0, index=spy.index)
    for i in range(1, len(spy)):
        if bool(entry_mask.iloc[i - 1]):
            out.iloc[i] = float(oc.iloc[i]) if np.isfinite(oc.iloc[i]) else 0.0
    return out


def overnight_5d_low(spy: pd.DataFrame) -> pd.Series:
    low5 = spy["close"] <= spy["close"].rolling(5, min_periods=5).min()
    return _overnight_one_night(spy, low5)


def overnight_3down(spy: pd.DataFrame) -> pd.Series:
    c2c = _spy_close_to_close(spy)
    three = (c2c < 0) & (c2c.shift(1) < 0) & (c2c.shift(2) < 0)
    return _overnight_one_night(spy, three)


def _mr_state_machine(
    spy: pd.DataFrame,
    entry: pd.Series,
    exit_sig: pd.Series,
    max_hold: int,
) -> pd.Series:
    """Signal at close t → position earns close-to-close returns from t+1 onward."""
    c2c = _spy_close_to_close(spy)
    pos = np.zeros(len(spy), dtype=bool)
    hold = 0
    for i in range(len(spy)):
        prev = pos[i - 1] if i else False
        if prev:
            hold += 1
            if bool(exit_sig.iloc[i - 1]) or hold > max_hold:
                pos[i] = False
                hold = 0
            else:
                pos[i] = True
        elif i > 0 and bool(entry.iloc[i - 1]):
            pos[i] = True
            hold = 1
    mask = pd.Series(pos, index=spy.index)
    return c2c.where(mask, 0.0)


def mr_rsi2_sma200(spy: pd.DataFrame) -> pd.Series:
    rsi2 = _rsi(spy["close"], 2)
    sma200 = spy["close"].rolling(200, min_periods=100).mean()
    entry = (rsi2 < 10) & (spy["close"] > sma200)
    exit_sig = rsi2 > 70
    return _mr_state_machine(spy, entry, exit_sig, max_hold=5)


def mr_ibs_sma200(spy: pd.DataFrame) -> pd.Series:
    ibs = _ibs(spy)
    sma200 = spy["close"].rolling(200, min_periods=100).mean()
    entry = (ibs < 0.2) & (spy["close"] > sma200)
    exit_sig = ibs > 0.8
    return _mr_state_machine(spy, entry, exit_sig, max_hold=3)


def _mom_12m(close: pd.Series) -> pd.Series:
    return close / close.shift(252) - 1.0


def dual_momentum_spy_tlt(spy: pd.DataFrame, tlt: pd.DataFrame) -> pd.Series:
    """Antonacci: pick SPY or TLT by 12m momentum; cash if absolute mom < 0."""
    idx = spy.index.intersection(tlt.index)
    spy_c = spy.loc[idx, "close"]
    tlt_c = tlt.loc[idx, "close"]
    m_spy = _mom_12m(spy_c)
    m_tlt = _mom_12m(tlt_c)
    spy_r = spy_c.pct_change().fillna(0.0)
    tlt_r = tlt_c.pct_change().fillna(0.0)
    # rebalance at month start
    month = pd.Series(idx, index=idx).dt.to_period("M")
    choice = pd.Series("cash", index=idx, dtype=object)
    for per, dates in month.groupby(month):
        d0 = dates.index[0]
        ms, mt = float(m_spy.loc[d0]), float(m_tlt.loc[d0])
        if not (np.isfinite(ms) and np.isfinite(mt)):
            pick = "cash"
        elif max(ms, mt) <= 0:
            pick = "cash"
        else:
            pick = "spy" if ms >= mt else "tlt"
        for d in dates.index:
            choice.loc[d] = pick
    ret = pd.Series(0.0, index=idx)
    for d in idx:
        c = choice.loc[d]
        if c == "spy":
            ret.loc[d] = float(spy_r.loc[d])
        elif c == "tlt":
            ret.loc[d] = float(tlt_r.loc[d])
    return ret


def rotation_spy_tlt_gld(
    spy: pd.DataFrame, tlt: pd.DataFrame, gld: pd.DataFrame
) -> pd.Series:
    """Monthly hold top-1 of SPY/TLT/GLD by 12-month momentum."""
    idx = spy.index.intersection(tlt.index).intersection(gld.index)
    panels = {"spy": spy.loc[idx, "close"], "tlt": tlt.loc[idx, "close"], "gld": gld.loc[idx, "close"]}
    rets = {k: v.pct_change().fillna(0.0) for k, v in panels.items()}
    moms = {k: _mom_12m(v) for k, v in panels.items()}
    month = pd.Series(idx, index=idx).dt.to_period("M")
    choice = pd.Series("spy", index=idx, dtype=object)
    for per, dates in month.groupby(month):
        d0 = dates.index[0]
        scores = {k: float(moms[k].loc[d0]) for k in panels}
        if not all(np.isfinite(v) for v in scores.values()):
            pick = "spy"
        else:
            pick = max(scores, key=scores.get)  # type: ignore[arg-type]
        for d in dates.index:
            choice.loc[d] = pick
    ret = pd.Series(0.0, index=idx)
    for d in idx:
        ret.loc[d] = float(rets[choice.loc[d]].loc[d])
    return ret


STRATEGIES = {
    "seasonality_turnaround_tuesday": lambda spy, tlt, gld: seasonality_turnaround_tuesday(spy),
    "seasonality_turn_of_month": lambda spy, tlt, gld: seasonality_turn_of_month(spy),
    "seasonality_opex_week": lambda spy, tlt, gld: seasonality_opex_week(spy),
    "overnight_close_to_open": lambda spy, tlt, gld: overnight_close_to_open(spy),
    "overnight_5d_low": lambda spy, tlt, gld: overnight_5d_low(spy),
    "overnight_3down": lambda spy, tlt, gld: overnight_3down(spy),
    "mr_rsi2_sma200": lambda spy, tlt, gld: mr_rsi2_sma200(spy),
    "mr_ibs_sma200": lambda spy, tlt, gld: mr_ibs_sma200(spy),
    "dual_momentum_spy_tlt": lambda spy, tlt, gld: dual_momentum_spy_tlt(spy, tlt),
    "rotation_spy_tlt_gld": lambda spy, tlt, gld: rotation_spy_tlt_gld(spy, tlt, gld),
}


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("--out-prefix", type=Path, default=DEFAULT_OUT)
    args = ap.parse_args()

    start = pd.Timestamp(args.start)
    end = pd.Timestamp(args.end) if str(args.end).strip() else None
    cap = float(args.capital)

    spy = _align(_fetch_ohlc("SPY", args.yahoo_period), start, end)
    tlt = _align(_fetch_ohlc("TLT", args.yahoo_period), start, end)
    gld = _align(_fetch_ohlc("GLD", args.yahoo_period), start, end)
    spy_bh = _spy_close_to_close(spy)

    results: dict[str, dict] = {}
    daily_panel: dict[str, pd.Series] = {"SPY_buy_hold": spy_bh}

    for name, fn in STRATEGIES.items():
        r = fn(spy, tlt, gld).reindex(spy.index).fillna(0.0)
        m = _metrics(r, cap)
        m["first_date"] = str(r.index.min().date()) if len(r) else ""
        m["last_date"] = str(r.index.max().date()) if len(r) else ""
        results[name] = m
        daily_panel[name] = r
        print(
            f"{name:32s}  ret={m['total_return_pct']:7.2f}%  "
            f"CAGR={m['cagr_pct']:6.2f}%  Sharpe={m['sharpe']:5.2f}  "
            f"DD={m['max_dd_pct']:6.2f}%  invested={m['pct_days_invested']:5.1f}%",
            flush=True,
        )

    spy_bh_m = _metrics(spy_bh, cap)
    results["SPY_buy_hold"] = spy_bh_m
    print(
        f"{'SPY_buy_hold':32s}  ret={spy_bh_m['total_return_pct']:7.2f}%  "
        f"CAGR={spy_bh_m['cagr_pct']:6.2f}%  Sharpe={spy_bh_m['sharpe']:5.2f}  "
        f"DD={spy_bh_m['max_dd_pct']:6.2f}%",
        flush=True,
    )

    # Correlation vs SPY and cross-sleeve
    corr_df = pd.DataFrame(daily_panel).dropna(how="all")
    corr = corr_df.corr().round(3)
    print("\nCorrelation vs SPY buy-hold:", flush=True)
    for name in STRATEGIES:
        print(f"  {name:32s}  rho={corr.loc[name, 'SPY_buy_hold']:.3f}", flush=True)

    # Rank by Sharpe
    ranked = sorted(
        [(k, v["sharpe"]) for k, v in results.items() if k != "SPY_buy_hold"],
        key=lambda x: x[1],
        reverse=True,
    )
    print("\nRanked by Sharpe:", flush=True)
    for i, (k, s) in enumerate(ranked, 1):
        print(f"  {i:2d}. {k} ({s:.3f})", flush=True)

    prefix = args.out_prefix.expanduser().resolve()
    prefix.parent.mkdir(parents=True, exist_ok=True)
    meta_path = Path(f"{prefix}_meta.json")
    corr_path = Path(f"{prefix}_correlation.csv")
    daily_path = Path(f"{prefix}_daily.csv")
    metrics_path = Path(f"{prefix}_metrics.txt")

    meta = {
        "command": " ".join(sys.argv),
        "start": str(start.date()),
        "end": str(end.date()) if end is not None else "",
        "capital": cap,
        "strategies": results,
        "ranked_by_sharpe": [{"name": k, "sharpe": s} for k, s in ranked],
    }
    meta_path.write_text(json.dumps(meta, indent=2))

    out_daily = pd.DataFrame(
        {k: v.values for k, v in daily_panel.items()},
        index=corr_df.index,
    )
    out_daily.index.name = "date"
    out_daily.to_csv(daily_path)
    corr.to_csv(corr_path)

    lines = [
        f"QS top ideas backtest  {start.date()} → {end.date() if end else 'latest'}  capital=${cap:,.0f}",
        "",
        f"{'strategy':32s} {'return%':>9s} {'CAGR%':>7s} {'Sharpe':>7s} {'maxDD%':>8s} {'inv%':>6s}",
    ]
    for name in list(STRATEGIES.keys()) + ["SPY_buy_hold"]:
        m = results[name]
        inv = m.get("pct_days_invested", 100.0)
        lines.append(
            f"{name:32s} {m['total_return_pct']:9.2f} {m['cagr_pct']:7.2f} "
            f"{m['sharpe']:7.3f} {m['max_dd_pct']:8.2f} {inv:6.1f}"
        )
    lines.extend(["", "Ranked by Sharpe:"])
    for i, (k, s) in enumerate(ranked, 1):
        lines.append(f"  {i}. {k} — {s:.3f}")
    metrics_path.write_text("\n".join(lines) + "\n")

    print(f"\nWrote {meta_path}", flush=True)
    print(f"Wrote {daily_path}", flush=True)


if __name__ == "__main__":
    main()
