#!/usr/bin/env python3
"""
Prototype **market-neutral sector L/S rotation** (SPDR 11 sectors).

Long top-k / short bottom-k by 12-minus-1 ``aqr_mom`` (monthly rebalance, daily PnL).
Dollar-neutral: +100% long leg, −100% short leg (~200% gross).

Example::

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

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 DataLoader
from RenTech.strategy_stack.main import _compute_daily_backtest_features, _load_sector_etf_dict
from RenTech.strategy_stack.multi_strategy_manager import (
    SPDR_SECTOR_TICKERS,
    SectorETFLongShort,
)

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


def _beta_vs_spy(r: pd.Series, spy_r: pd.Series) -> tuple[float, float]:
    aligned = pd.DataFrame({"s": r, "spy": spy_r.reindex(r.index).fillna(0.0)}).dropna()
    if len(aligned) < 10:
        return float("nan"), float("nan")
    rho = float(aligned.corr().iloc[0, 1])
    sv = float(aligned["spy"].var())
    beta = float(aligned["s"].cov(aligned["spy"]) / sv) if sv > 1e-14 else float("nan")
    return beta, rho


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("--yahoo-period", default="max")
    ap.add_argument("--top-k", type=int, default=3, help="Sectors per leg (default 3)")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT_PREFIX)
    ap.add_argument("--cash-yield", type=float, default=0.04)
    args = ap.parse_args()

    etf_dict = _load_sector_etf_dict(args.yahoo_period)
    spy_df = _compute_daily_backtest_features(
        DataLoader().fetch_daily("SPY", period=args.yahoo_period)
    )
    eng = SectorETFLongShort()
    rebal = eng.generate_rebalance_log(etf_dict, top_k=int(args.top_k))
    daily_ret = eng.generate_returns(
        etf_dict,
        top_k=int(args.top_k),
        cash_annual_yield=float(args.cash_yield),
        verbose=True,
    )

    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_usd = cap * (1.0 + r).cumprod()
    n = len(r)
    years = n / 252.0
    end_eq = float(eq_usd.iloc[-1])
    tot = end_eq / cap - 1.0
    cagr = (end_eq / cap) ** (1.0 / years) - 1.0 if years > 0 else float("nan")
    max_dd = float((eq_usd / eq_usd.cummax() - 1.0).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")

    spy_r = spy_df["ret"].astype(float)
    spy_r.index = pd.to_datetime(spy_r.index).tz_localize(None)
    beta, rho = _beta_vs_spy(r, spy_r)

    gross = 2.0  # 100% long + 100% short
    margin_usd = eq_usd * 0.25 * gross

    prefix = args.out_prefix.expanduser().resolve()
    prefix.parent.mkdir(parents=True, exist_ok=True)
    daily_path = Path(f"{prefix}_daily.csv")
    rebal_path = Path(f"{prefix}_rebalances.csv")
    meta_path = Path(f"{prefix}_meta.json")

    pd.DataFrame(
        {
            "date": r.index.strftime("%Y-%m-%d"),
            "daily_ret": r.values,
            "daily_pnl_usd": pnl.values,
            "equity_usd": eq_usd.values,
            "margin_usd": margin_usd.values,
        }
    ).to_csv(daily_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:
        rebal.to_csv(rebal_path, index=False)

    cmd = (
        f"cd {_REPO} && PYTHONUNBUFFERED=1 .venv/bin/python "
        f"RenTech/strategy_stack/run_sector_ls_momentum_standard.py "
        f"--start {args.start} --end {args.end} --top-k {args.top_k} --capital {cap:.0f}"
    )
    meta = {
        "strategy": "sector_ls_momentum",
        "token": "sector_ls",
        "market_neutral": True,
        "description": "SPDR 11 sectors: long top-k / short bottom-k by aqr_mom",
        "tickers_universe": SPDR_SECTOR_TICKERS,
        "top_k_per_leg": int(args.top_k),
        "capital_usd": cap,
        "start": str(r.index.min().date()),
        "end": str(r.index.max().date()),
        "n_trading_days": int(n),
        "total_return_pct": round(tot * 100.0, 2),
        "cagr_pct": round(cagr * 100.0, 2),
        "sharpe_daily": round(sharpe, 3),
        "max_drawdown_pct": round(max_dd * 100.0, 2),
        "beta_vs_spy": round(beta, 3),
        "corr_vs_spy": round(rho, 3),
        "command": cmd,
        "daily_csv": str(daily_path),
    }
    meta_path.write_text(json.dumps(meta, indent=2) + "\n", encoding="utf-8")

    print(
        f"\nSector L/S momentum: return {meta['total_return_pct']:+.1f}%  "
        f"Sharpe {sharpe:.2f}  maxDD {meta['max_drawdown_pct']:.1f}%  "
        f"β(SPY) {beta:.2f}  ρ(SPY) {rho:.2f}",
        flush=True,
    )
    print(f"Wrote {daily_path}", flush=True)


if __name__ == "__main__":
    main()
