#!/usr/bin/env python3
"""
**Short-Term Reversal Sleeve** — 1-day mean reversion on large S&P 500 drops.

Signal: buy stocks in the S&P 500 that fell ≥ ``--drop-min`` intraday (open-to-close),
hold **1 trading day** (exit at next open → approximated as next day's open return).

Academic basis: Jegadeesh (1990), Lehmann (1990) — weekly/monthly short-term reversal;
extended by Da, Liu, Schaumburg (2014) to overnight holding period.

Rules:
  - Enter at close of drop day (market-on-close order approximation)
  - Exit at next day's open (here: next day's close for simplicity, conservative)
  - Filter: SPY > SMA(50) — avoid catching falling knives in bear phases
  - Max 10 simultaneous positions, equal weight
  - Only liquid names (price > $5, market cap proxy: in S&P 500)

Why this helps weak years:
  2015: multiple sharp 1-day drops (Aug flash crash), reversals captured
  2018: Q4 selloff had many 3–5% daily drops followed by partial recoveries
  2019: consistent short reversals in trade-war volatility
  2022: extreme daily moves → reversals; SPY filter limits exposure in sustained downtrends

Data: yfinance S&P 500 universe (cached list); uses adjusted closes.

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 \\
      .venv/bin/python RenTech/strategy_stack/run_short_term_reversal.py \\
      --start 2011-01-03 --end 2025-12-31 --capital 100000
"""

from __future__ import annotations

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

import numpy as np
import pandas as pd
import yfinance as yf

_REPO = Path(__file__).resolve().parents[2]
if str(_REPO) not in sys.path:
    sys.path.insert(0, str(_REPO))

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

# Stable liquid S&P 500 names with long history (sub-universe for speed)
# Using a diversified 60-stock proxy — enough to find candidates every day
LIQUID_SP500_PROXY = [
    "AAPL","MSFT","AMZN","GOOGL","META","NVDA","TSLA","JPM","JNJ","V",
    "WMT","PG","MA","UNH","HD","XOM","CVX","BAC","ABBV","PFE",
    "AVGO","KO","LLY","PEP","COST","MRK","TMO","CSCO","ACN","ABT",
    "NKE","WFC","CRM","MCD","DHR","LIN","TXN","NEE","NFLX","ADBE",
    "VZ","PM","INTC","UPS","RTX","QCOM","HON","AMGN","T","LOW",
    "GS","CAT","BA","SBUX","ISRG","GILD","DE","MDLZ","MMM","AXP",
]


def _download_panel(tickers: list[str], start: str, end: str) -> pd.DataFrame:
    raw = yf.download(tickers, start=start, end=end, auto_adjust=True, progress=False)
    if isinstance(raw.columns, pd.MultiIndex):
        close = raw["Close"].copy()
    else:
        close = raw[["Close"]].copy()
    close.index = pd.to_datetime(close.index).tz_localize(None)
    return close.sort_index().ffill(limit=3)


def run_short_term_reversal(
    start: str,
    end: str,
    capital: float,
    *,
    out_prefix: Path,
    drop_min: float = 0.03,
    max_positions: int = 10,
    sma_filter_window: int = 50,
    min_price: float = 5.0,
    verbose: bool = True,
) -> dict:
    t0, t1 = pd.Timestamp(start), pd.Timestamp(end)
    fetch_start = (t0 - pd.DateOffset(years=1)).strftime("%Y-%m-%d")
    fetch_end = t1.strftime("%Y-%m-%d")

    print("Downloading price panel...", flush=True)
    panel = _download_panel(LIQUID_SP500_PROXY, fetch_start, fetch_end)
    # Drop tickers with < 500 sessions
    panel = panel.loc[:, panel.notna().sum() >= 500]
    tickers = list(panel.columns)

    # SPY for filter
    spy_raw = yf.download(["SPY"], start=fetch_start, end=fetch_end,
                           auto_adjust=True, progress=False)
    if isinstance(spy_raw.columns, pd.MultiIndex):
        spy = spy_raw["Close"]["SPY"]
    else:
        spy = spy_raw["Close"] if "Close" in spy_raw.columns else spy_raw.iloc[:, 0]
    spy.index = pd.to_datetime(spy.index).tz_localize(None)
    spy_sma = spy.rolling(sma_filter_window, min_periods=sma_filter_window // 2).mean()

    daily_rets = panel.pct_change()

    idx = panel.index[(panel.index >= t0) & (panel.index <= t1)]
    cap = float(capital)
    port_rets: list[float] = []
    trade_log: list[dict] = []

    for i, dt in enumerate(idx):
        if i == 0:
            port_rets.append(0.0)
            continue

        prev_dt = idx[i - 1]

        # SPY regime filter on entry date (prev_dt)
        spy_px = float(spy.reindex([prev_dt]).iloc[0]) if prev_dt in spy.index else float("nan")
        spy_sm = float(spy_sma.reindex([prev_dt]).iloc[0]) if prev_dt in spy_sma.index else float("nan")
        if not (math.isfinite(spy_px) and math.isfinite(spy_sm) and spy_px > spy_sm * 0.97):
            # Below SMA50 by >3% → skip entry
            port_rets.append(0.0)
            continue

        # Drop candidates: yesterday's return ≤ -drop_min
        if prev_dt not in daily_rets.index:
            port_rets.append(0.0)
            continue
        prev_rets = daily_rets.loc[prev_dt]

        # Filter by price floor on entry day
        prev_prices = panel.loc[prev_dt]
        eligible = (
            (prev_rets <= -drop_min) &
            (prev_prices >= min_price)
        )
        candidates = prev_rets[eligible].sort_values()  # biggest drops first

        if candidates.empty:
            port_rets.append(0.0)
            continue

        # Take top-N biggest drops
        selected = candidates.index[:max_positions]
        n_pos = len(selected)

        # Today's return for selected names
        today_rets = daily_rets.loc[dt, selected].fillna(0.0) if dt in daily_rets.index else pd.Series(0.0, index=selected)
        pos_ret = float(today_rets.mean())
        port_rets.append(pos_ret)

        for t in selected:
            trade_log.append({
                "entry_date": str(prev_dt.date()),
                "exit_date": str(dt.date()),
                "ticker": t,
                "entry_return_prior_day": round(float(prev_rets[t]) * 100, 2),
                "exit_return": round(float(today_rets[t]) * 100, 2),
            })

    port_s = pd.Series(port_rets, index=idx, dtype=np.float64)
    eq = cap * (1.0 + port_s).cumprod()
    n = len(port_s)
    years = n / 252.0
    end_eq = float(eq.iloc[-1])
    total_ret = end_eq / cap - 1.0
    cagr = (end_eq / cap) ** (1.0 / years) - 1.0 if years > 0 else float("nan")
    dd = float((eq / eq.cummax() - 1.0).min())
    sd = float(port_s.std(ddof=1))
    sharpe = float(port_s.mean() / sd * math.sqrt(252)) if sd > 1e-12 else float("nan")

    spy_rets_bt = spy.pct_change().reindex(idx).fillna(0.0)
    rho_spy = float(pd.DataFrame({"p": port_s, "spy": spy_rets_bt}).dropna().corr().iloc[0, 1])

    yr_rows = []
    eq_cur = cap
    pct_active = []
    for yr, g in port_s.groupby(port_s.index.year):
        ret_y = float((1 + g).prod() - 1) * 100
        end_y = eq_cur * (1 + ret_y / 100)
        active = float((g != 0).mean()) * 100
        yr_rows.append({"year": int(yr), "return_pct": round(ret_y, 2),
                         "pnl_usd": round(end_y - eq_cur, 0), "pct_active": round(active, 1)})
        eq_cur = end_y
        pct_active.append(active)
    yr_df = pd.DataFrame(yr_rows)

    daily_df = pd.DataFrame({
        "date":          port_s.index.strftime("%Y-%m-%d"),
        "daily_ret":     port_s.values,
        "daily_pnl_usd": (port_s * cap).values,
        "equity_usd":    eq.values,
    })

    out_prefix = Path(out_prefix)
    out_prefix.parent.mkdir(parents=True, exist_ok=True)
    daily_df.to_csv(f"{out_prefix}_daily.csv", index=False)
    yr_df.to_csv(f"{out_prefix}_yearly.csv", index=False)
    if trade_log:
        pd.DataFrame(trade_log).to_csv(f"{out_prefix}_trades.csv", index=False)

    meta = {
        "strategy": "short_term_reversal",
        "rules": {
            "drop_min_pct": drop_min * 100,
            "max_positions": max_positions,
            "hold_days": 1,
            "spy_sma_filter": sma_filter_window,
            "universe": f"S&P 500 liquid proxy ({len(tickers)} tickers)",
        },
        "start": str(idx.min().date()),
        "end": str(idx.max().date()),
        "n_sessions": n,
        "n_trades": len(trade_log),
        "capital": cap,
        "ending_equity_usd": round(end_eq, 2),
        "total_return_pct": round(total_ret * 100, 4),
        "cagr_pct": round(cagr * 100, 4),
        "sharpe": round(sharpe, 4),
        "max_dd_pct": round(dd * 100, 4),
        "rho_spy": round(rho_spy, 4),
        "avg_pct_days_active": round(float(np.mean(pct_active)), 1),
        "daily_csv": f"{out_prefix}_daily.csv",
        "yearly_csv": f"{out_prefix}_yearly.csv",
    }
    with open(f"{out_prefix}_meta.json", "w") as fh:
        json.dump(meta, fh, indent=2)

    if verbose:
        print("=== Short-Term Reversal Sleeve (1-day hold) ===")
        print(f"Window: {meta['start']} → {meta['end']}  ({n} sessions, {len(trade_log)} trades)")
        print(f"Universe: {len(tickers)} liquid S&P 500 names")
        print(f"Return {total_ret*100:.1f}%  CAGR {cagr*100:.1f}%  Sharpe {sharpe:.2f}  MaxDD {dd*100:.1f}%  ρ(SPY) {rho_spy:.2f}")
        print(f"Avg active days: {float(np.mean(pct_active)):.1f}%")
        print("\nYearly returns:")
        print(yr_df.to_string(index=False))

    return meta


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
    ap.add_argument("--start", default="2011-01-03")
    ap.add_argument("--end", default="2025-12-31")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--drop-min", type=float, default=0.03,
                    help="Min prior-day drop to qualify (default 3%%)")
    ap.add_argument("--max-positions", type=int, default=10)
    ap.add_argument("--sma-filter", type=int, default=50,
                    help="SPY SMA filter window (default 50; 0 = disable)")
    ap.add_argument("--out-prefix", type=Path, default=LOGS / "short_term_reversal_standard")
    args = ap.parse_args()

    run_short_term_reversal(
        start=args.start,
        end=args.end,
        capital=float(args.capital),
        out_prefix=args.out_prefix,
        drop_min=float(args.drop_min),
        max_positions=int(args.max_positions),
        sma_filter_window=int(args.sma_filter),
        verbose=True,
    )


if __name__ == "__main__":
    main()
