#!/usr/bin/env python3
"""
**S&P 500 MA slope momentum** — monthly top-N by dual EMA slope rank score.

Ranks constituents on stage-4 SPY logic (EMA10 + EMA50 slopes both positive,
close > fast EMA). Default score: ``fast_slope × slow_slope``; holds top 10
equal-weight, rebalanced month-end. **Default stop:** 2× ATR chandelier between
rebalances (per-name exit to cash until next rebalance).

Example::

    cd /Users/robzingale/trading_bot
    PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_ma_slope_sp500_topn_standard.py \\
        --start 2016-01-04 --top-n 10

Smoke (50 names)::

    PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_ma_slope_sp500_topn_standard.py \\
        --start 2016-01-04 --top-n 10 --max-tickers 50
"""

from __future__ import annotations

import argparse
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.equity_universe_loaders import load_equity_panel_dict
from RenTech.strategy_stack.ma_slope_cross_sectional import (
    MaSlopeCrossSectional,
    MaSlopeCrossSectionalConfig,
)
from RenTech.strategy_stack.main import _compute_daily_backtest_features

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


def _stop_slug(cfg: MaSlopeCrossSectionalConfig) -> str:
    if cfg.stop_mode == "none":
        return ""
    if cfg.stop_mode == "atr_trail":
        m = cfg.atr_multiplier
        label = str(int(m)) if float(m).is_integer() else f"{m:g}"
        return f"_atr{label}x"
    if cfg.stop_mode == "pct_trail":
        return f"_pct{int(cfg.pct_trail_stop * 100)}"
    return f"_{cfg.stop_mode}"


def _yearly_stats(r: pd.Series) -> pd.DataFrame:
    rows: list[dict] = []
    for year, g in r.groupby(r.index.year):
        eq = (1.0 + g).cumprod()
        ret = float(eq.iloc[-1] - 1.0) * 100.0
        dd = float((eq / eq.cummax() - 1.0).min()) * 100.0
        rows.append({"year": int(year), "return_pct": ret, "max_dd_pct": dd, "n_days": len(g)})
    return pd.DataFrame(rows)


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("--yahoo-period", default="max")
    ap.add_argument("--top-n", type=int, default=10, help="Names held each rebalance (default 10)")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--max-tickers", type=int, default=0, help="Cap universe for smoke tests (0=all)")
    ap.add_argument("--refresh-cache", action="store_true")
    ap.add_argument("--cash-yield", type=float, default=0.04)
    ap.add_argument("--rank-metric", default="dual_product", choices=(
        "dual_product", "dual_blend", "fast", "slow", "dual_min",
    ))
    ap.add_argument("--rebalance", default="monthly", choices=("monthly", "weekly"))
    ap.add_argument("--fast-period", type=int, default=10)
    ap.add_argument("--slow-period", type=int, default=50)
    ap.add_argument("--fast-lookback", type=int, default=10)
    ap.add_argument("--slow-lookback", type=int, default=5)
    ap.add_argument("--entry-slope-min", type=float, default=0.0)
    ap.add_argument("--no-price-filter", action="store_true")
    ap.add_argument(
        "--no-stops",
        action="store_true",
        help="Disable per-name stops (default: 2× ATR trail between rebalances)",
    )
    ap.add_argument(
        "--stop-mode",
        default="",
        choices=("", "none", "atr_trail", "pct_trail", "portfolio_dd"),
        help="Override stop mode (default: atr_trail)",
    )
    ap.add_argument("--atr-multiplier", type=float, default=2.0)
    ap.add_argument("--atr-period", type=int, default=14)
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT)
    args = ap.parse_args()

    equity_dict = load_equity_panel_dict(
        "sp500",
        args.yahoo_period,
        max_tickers=int(args.max_tickers),
        refresh_cache=bool(args.refresh_cache),
    )
    if not equity_dict:
        raise SystemExit("Empty S&P 500 equity panel")

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

    stop_mode = "none" if args.no_stops else (args.stop_mode or "atr_trail")
    cfg = MaSlopeCrossSectionalConfig(
        fast_period=int(args.fast_period),
        slow_period=int(args.slow_period),
        fast_lookback=int(args.fast_lookback),
        slow_lookback=int(args.slow_lookback),
        entry_slope_min=float(args.entry_slope_min),
        price_above_ma=not args.no_price_filter,
        rank_metric=args.rank_metric,
        rebalance=args.rebalance,
        cash_annual_yield=float(args.cash_yield),
        stop_mode=stop_mode,
        atr_period=int(args.atr_period),
        atr_multiplier=float(args.atr_multiplier),
    )
    eng = MaSlopeCrossSectional(config=cfg)

    print(f"Building slope panels for {len(equity_dict)} names …", flush=True)
    daily_ret = eng.generate_returns(equity_dict, top_n=int(args.top_n), verbose=True)
    rebal = eng.generate_rebalance_log(equity_dict, top_n=int(args.top_n))

    daily_ret = daily_ret.sort_index()
    daily_ret.index = pd.to_datetime(daily_ret.index).tz_localize(None)
    mask = daily_ret.index >= pd.Timestamp(args.start)
    if args.end.strip():
        mask &= daily_ret.index <= pd.Timestamp(args.end)
    r = daily_ret.loc[mask].fillna(0.0).astype(np.float64)

    cap = float(args.capital)
    pnl = r * cap
    eq_unit = (1.0 + r).cumprod()
    eq_usd = cap * eq_unit

    n = len(r)
    years = n / 252.0
    end_eq = float(eq_usd.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 = eq_usd / eq_usd.cummax() - 1.0
    max_dd = float(dd.min())
    sd = float(r.std(ddof=1)) if n > 1 else float("nan")
    sharpe = float(r.mean() / sd * np.sqrt(252.0)) if sd > 1e-12 else float("nan")
    vol_ann = sd * np.sqrt(252.0) if np.isfinite(sd) else float("nan")

    spy_r = spy_df["ret"].astype(float)
    spy_r.index = pd.to_datetime(spy_r.index).tz_localize(None)
    rho_spy = float(pd.DataFrame({"mom": r, "spy": spy_r}).dropna().corr().iloc[0, 1])

    prefix = args.out_prefix.expanduser().resolve()
    prefix.parent.mkdir(parents=True, exist_ok=True)
    slug = f"top{int(args.top_n)}_{args.rank_metric}_{args.rebalance}{_stop_slug(cfg)}"
    daily_path = Path(f"{prefix}_{slug}_daily.csv")
    rebal_path = Path(f"{prefix}_{slug}_rebalances.csv")
    yearly_path = Path(f"{prefix}_{slug}_yearly.csv")
    meta_path = Path(f"{prefix}_{slug}_meta.json")
    metrics_path = Path(f"{prefix}_{slug}_metrics.txt")

    out = pd.DataFrame(
        {
            "date": r.index.strftime("%Y-%m-%d"),
            "daily_ret": r.values,
            "daily_pnl_usd": pnl.values,
            "equity_unit": eq_unit.values,
            "equity_usd": eq_usd.values,
        }
    )
    out.to_csv(daily_path, index=False)

    yearly = _yearly_stats(r)
    yearly.to_csv(yearly_path, index=False)

    if len(rebal):
        rmask = pd.to_datetime(rebal["effective_date"]) >= pd.Timestamp(args.start)
        if args.end.strip():
            rmask &= pd.to_datetime(rebal["effective_date"]) <= pd.Timestamp(args.end)
        rebal.loc[rmask].to_csv(rebal_path, index=False)
    else:
        pd.DataFrame().to_csv(rebal_path, index=False)

    cmd = (
        f"cd {_REPO} && PYTHONUNBUFFERED=1 .venv/bin/python "
        f"RenTech/strategy_stack/run_ma_slope_sp500_topn_standard.py "
        f"--start {args.start} --top-n {args.top_n} --rank-metric {args.rank_metric} "
        f"--yahoo-period {args.yahoo_period}"
    )
    if args.max_tickers:
        cmd += f" --max-tickers {args.max_tickers}"
    if args.no_stops:
        cmd += " --no-stops"
    elif args.stop_mode:
        cmd += f" --stop-mode {args.stop_mode}"
    if cfg.stop_mode == "atr_trail" and cfg.atr_multiplier != 2.0:
        cmd += f" --atr-multiplier {cfg.atr_multiplier:g}"

    stop_line = "none"
    if cfg.stop_mode == "atr_trail":
        stop_line = f"ATR {cfg.atr_multiplier:g}× trail (period {cfg.atr_period})"
    elif cfg.stop_mode != "none":
        stop_line = cfg.stop_mode

    meta = {
        "strategy": "MaSlopeCrossSectional",
        "token": "ma_slope_sp500_topn",
        "universe": "sp500",
        "n_universe": len(equity_dict),
        "top_n": int(args.top_n),
        "config": asdict(cfg),
        "capital_usd": cap,
        "start": str(r.index.min().date()),
        "end": str(r.index.max().date()),
        "n_trading_days": int(n),
        "total_return_pct": round(100.0 * total_ret, 4),
        "cagr_pct": round(100.0 * cagr, 4),
        "sharpe_daily": round(sharpe, 4),
        "vol_annual_pct": round(100.0 * vol_ann, 4),
        "max_drawdown_pct": round(100.0 * max_dd, 4),
        "corr_vs_spy_daily": round(rho_spy, 4),
        "command": cmd,
        "daily_csv": str(daily_path),
        "yearly_csv": str(yearly_path),
        "rebalances_csv": str(rebal_path),
        "n_rebalance_rows": int(len(rebal)),
    }
    meta_path.write_text(json.dumps(meta, indent=2) + "\n")

    metrics_path.write_text(
        f"""# MA slope S&P 500 top-{args.top_n} ({meta['start']} -> {meta['end']})

Command:
{cmd}

Rank score: {args.rank_metric}  |  Rebalance: {args.rebalance}
Stops: {stop_line}
Fast EMA({args.fast_period}) slope LB {args.fast_lookback}
Slow EMA({args.slow_period}) slope LB {args.slow_lookback}
Eligibility: both slopes > {args.entry_slope_min}, price > fast EMA

Headline (${cap:,.0f} notional, {len(equity_dict)} names):
  Total return: {meta['total_return_pct']:.2f}%
  CAGR: {meta['cagr_pct']:.2f}%
  Sharpe: {meta['sharpe_daily']:.3f}
  Max DD: {meta['max_drawdown_pct']:.2f}%
  Corr vs SPY: {meta['corr_vs_spy_daily']:.3f}

Artifacts:
  {daily_path}
  {yearly_path}
  {rebal_path}
  {meta_path}
"""
    )

    print(metrics_path.read_text(), flush=True)
    print(f"Wrote {daily_path}", flush=True)
    print(f"Wrote {rebal_path}  ({len(rebal)} rows)", flush=True)


if __name__ == "__main__":
    main()
