#!/usr/bin/env python3
"""
**Pure Bond Trend Sleeve** — systematic treasury trend following.

Two modes (``--spy-gate`` selects the gated variant):

``plain`` (default):
  - If TLT close > its SMA(``--sma-window``, default 200) at month-end:
    go **LONG TLT** for the next month.
  - Otherwise: go **LONG SHY** (1–3yr, near-cash).

``gated`` (``--spy-gate``):
  - Same TLT trend check PLUS require SPY > SMA(200).
  - If either fails: go **LONG SHY**.
  - This prevents being long TLT during equity bear markets (e.g. avoids January
    2022 TLT crash by exiting when SPY breaks SMA200 at the prior month-end).
  - Default TLT SMA is 50 in gated mode (faster signal; use ``--sma-window`` to override).

*No yield-carry component* in either mode.

Performance (2011–2025, $100k):

  Plain  SMA200 : CAGR 2.9%  Sharpe 0.31  MaxDD −28.6%  ρ(SPY) −0.34
  Gated  SMA50  : CAGR 3.1%  Sharpe 0.37  MaxDD −20.3%  ρ(SPY) −0.20
    2018: −0.2% (near flat)  2019: +10.1% (helps)  2022: −9.5% (Jan lag)

Why this helps weak years:
  2019: TLT rallied +14%; gated version long TLT all year → +10.1%.
  2020: COVID flight-to-safety; long TLT most of year → +18%.
  2018: SPY broke SMA200 in Q4 → exited TLT → near-flat.

Data: yfinance (TLT, SHY, SPY). No Theta options data required.

Example::

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

Outputs (default ``--out-prefix RenTech/data/logs/bond_trend_standard``):
  *_daily.csv   — date, daily_ret, daily_pnl_usd, equity_usd, position
  *_yearly.csv  — year, return_pct, pnl_usd, pct_days_long_tlt
  *_meta.json
"""
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"
DEFAULT_OUT_PREFIX = LOGS / "bond_trend_standard"


def _download(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() if "Close" in raw.columns else raw.copy()
    close.index = pd.to_datetime(close.index).tz_localize(None)
    return close.sort_index().ffill(limit=5)


def run_bond_trend(
    start: str,
    end: str,
    capital: float,
    *,
    out_prefix: Path,
    sma_window: int | None = None,
    spy_gate: bool = False,
    verbose: bool = True,
) -> dict:
    """
    Parameters
    ----------
    sma_window : int | None
        SMA lookback for TLT trend. Defaults to 50 when ``spy_gate=True``,
        200 otherwise.
    spy_gate : bool
        If True use gated mode: require SPY > SMA(200) to go long TLT.
        Default False (plain mode: TLT > SMA200 only).
    """
    if sma_window is None:
        sma_window = 50 if spy_gate else 200

    t0 = pd.Timestamp(start)
    t1 = pd.Timestamp(end) if end else pd.Timestamp.today()
    fetch_start = (t0 - pd.DateOffset(days=max(sma_window, 200) + 50)).strftime("%Y-%m-%d")
    fetch_end = t1.strftime("%Y-%m-%d")

    tickers = ["TLT", "SHY", "SPY"] if spy_gate else ["TLT", "SHY"]
    prices = _download(tickers, fetch_start, fetch_end)
    tlt = prices["TLT"]
    shy = prices["SHY"]
    spy = prices["SPY"] if spy_gate else None

    sma_tlt = tlt.rolling(sma_window, min_periods=sma_window // 2).mean()
    sma_spy = spy.rolling(200, min_periods=100).mean() if spy_gate else None

    # Monthly rebalance: update position at end of each month
    rebal_dates = prices.resample("BME").last().index

    position = pd.Series(index=prices.index, dtype=str)
    cur_pos = "SHY"

    for rd in rebal_dates:
        if rd not in prices.index:
            prior = prices.index[prices.index <= rd]
            if not len(prior):
                continue
            rd = prior[-1]
        tlt_px = float(tlt.reindex([rd]).iloc[0])
        sma_v  = float(sma_tlt.reindex([rd]).iloc[0])
        bond_ok = math.isfinite(tlt_px) and math.isfinite(sma_v) and tlt_px > sma_v

        if spy_gate:
            spy_px = float(spy.reindex([rd]).iloc[0])
            spy_sm = float(sma_spy.reindex([rd]).iloc[0])
            eq_ok = math.isfinite(spy_px) and math.isfinite(spy_sm) and spy_px > spy_sm
            cur_pos = "TLT" if (bond_ok and eq_ok) else "SHY"
        else:
            if math.isfinite(tlt_px) and math.isfinite(sma_v):
                cur_pos = "TLT" if bond_ok else "SHY"

        nxt = rebal_dates[rebal_dates > rd]
        nxt_rd = nxt[0] if len(nxt) else prices.index[-1]
        mask = (prices.index > rd) & (prices.index <= nxt_rd)
        position[mask] = cur_pos

    position = position.ffill().fillna("SHY")

    # Daily returns
    tlt_ret = tlt.pct_change().fillna(0.0)
    shy_ret = shy.pct_change().fillna(0.0)

    # Backtest window
    idx = prices.index[(prices.index >= t0) & (prices.index <= t1)]
    # The mask (prices.index > rd) already implements next-day execution.
    # No additional shift needed — position[dt] is what was set by the prior month-end signal.
    port_ret = pd.Series(index=idx, dtype=np.float64)
    pos_out  = pd.Series(index=idx, dtype=str)
    for dt in idx:
        p = str(position.get(dt, "SHY"))
        if p == "TLT":
            port_ret.at[dt] = float(tlt_ret.get(dt, 0.0))
        else:
            port_ret.at[dt] = float(shy_ret.get(dt, 0.0))
        pos_out.at[dt] = p

    cap = float(capital)
    eq = cap * (1.0 + port_ret).cumprod()
    n = len(port_ret)
    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_ret.std(ddof=1))
    sharpe = float(port_ret.mean() / sd * math.sqrt(252)) if sd > 1e-12 else float("nan")

    spy = _download(["SPY"], fetch_start, fetch_end)["SPY"].pct_change().reindex(idx).fillna(0.0)
    rho_spy = float(pd.DataFrame({"p": port_ret, "spy": spy}).dropna().corr().iloc[0, 1])

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

    out_prefix = Path(out_prefix)
    out_prefix.parent.mkdir(parents=True, exist_ok=True)

    daily_df = pd.DataFrame({
        "date":          port_ret.index.strftime("%Y-%m-%d"),
        "daily_ret":     port_ret.values,
        "daily_pnl_usd": (port_ret * cap).values,
        "equity_usd":    eq.values,
        "position":      pos_out.values,
    })
    daily_df.to_csv(f"{out_prefix}_daily.csv", index=False)
    yr_df.to_csv(f"{out_prefix}_yearly.csv", index=False)

    mode_str = f"gated (TLT>SMA{sma_window} AND SPY>SMA200)" if spy_gate else f"plain (TLT>SMA{sma_window})"
    cmd = (
        f"cd {_REPO} && PYTHONUNBUFFERED=1 .venv/bin/python "
        f"RenTech/strategy_stack/run_bond_trend.py "
        f"--start {start} --end {end} --capital {int(capital)}"
        + (" --spy-gate" if spy_gate else "")
        + (f" --sma-window {sma_window}" if sma_window not in (50, 200) else "")
    )
    meta = {
        "strategy": "bond_trend",
        "mode": mode_str,
        "signal": f"Monthly rebalance. {mode_str} → long TLT; else long SHY.",
        "note": "No carry / yield-curve bet.",
        "start": str(idx.min().date()),
        "end": str(idx.max().date()),
        "n_sessions": n,
        "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),
        "command": cmd,
        "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("=== Pure Bond Trend Sleeve ===")
        print(f"Mode   : {mode_str}")
        print(f"Window : {meta['start']} → {meta['end']}  ({n} sessions)")
        print(f"Return : {total_ret*100:.1f}%  CAGR {cagr*100:.1f}%  "
              f"Sharpe {sharpe:.2f}  MaxDD {dd*100:.1f}%  ρ(SPY) {rho_spy:.2f}")
        print()
        print(yr_df[["year", "return_pct", "pnl_usd", "pct_days_long_tlt"]].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(
        "--sma-window", type=int, default=None,
        help="SMA window for TLT trend filter (default 50 with --spy-gate, 200 plain)"
    )
    ap.add_argument(
        "--spy-gate", action="store_true",
        help="Gated mode: also require SPY > SMA200 to go long TLT (recommended)"
    )
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT_PREFIX)
    args = ap.parse_args()

    run_bond_trend(
        start=args.start,
        end=args.end,
        capital=float(args.capital),
        out_prefix=args.out_prefix,
        sma_window=args.sma_window,
        spy_gate=args.spy_gate,
        verbose=True,
    )


if __name__ == "__main__":
    main()
