#!/usr/bin/env python3
"""
Export Markov VIX calendar-spread backtest: equity curve, trade log, HTML report.

Reads ``markov_vix_spread_{daily,signals}.csv`` (from
``run_markov_vix_spread_backtest.py``) and writes:

  * ``{prefix}_equity_curve.csv``  — daily equity, drawdown, cumulative return
  * ``{prefix}_trade_log.csv``     — one row per spread position (entry → exit)
  * ``{prefix}_report.html``       — equity chart + headline metrics + trade table
  * ``{prefix}_equity_curve.png``  — chart image embedded in HTML

Example::

    cd /Users/robzingale/trading_bot
    PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/export_markov_vix_spread_report.py
"""

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))

LOGS = _REPO / "RenTech" / "data" / "logs"
DEFAULT_PREFIX = LOGS / "markov_vix_spread"
VIX_FUT_PATH = _REPO / "RenTech" / "data" / "vix_futures_cboe.parquet"

SPREAD_NEAR_FAR: dict[str, tuple[str, str]] = {
    "M1_M2": ("vx1", "vx2"),
    "M2_M3": ("vx2", "vx3"),
    "M3_M4": ("vx3", "vx4"),
    "M4_M5": ("vx4", "vx5"),
}

# Near-month expiry column used to detect contract rolls (constant-tenor series).
SPREAD_NEAR_EXPIRY: dict[str, str] = {
    "M1_M2": "vx1_expiry",
    "M2_M3": "vx2_expiry",
    "M3_M4": "vx3_expiry",
    "M4_M5": "vx4_expiry",
}


def _spread_daily_returns(panel: pd.DataFrame, master: pd.DatetimeIndex) -> pd.DataFrame:
    out = pd.DataFrame(index=master, dtype=np.float64)
    for label, (near, far) in SPREAD_NEAR_FAR.items():
        if near not in panel.columns or far not in panel.columns:
            continue
        n = panel[near].astype(float).reindex(master).ffill()
        f = panel[far].astype(float).reindex(master).ffill()
        out[label] = n.pct_change().fillna(0.0) - f.pct_change().fillna(0.0)
    return out


def _near_expiry_series(panel: pd.DataFrame, spread: str, dates: pd.Series) -> pd.Series:
    col = SPREAD_NEAR_EXPIRY.get(spread)
    if col is None or col not in panel.columns:
        return pd.Series(pd.NaT, index=dates.index)
    ex = pd.to_datetime(panel[col], errors="coerce")
    ex.index = pd.to_datetime(panel.index).tz_localize(None)
    return ex.reindex(pd.to_datetime(dates)).ffill()


def _close_trade_row(
    *,
    trade_id: int,
    spread: str,
    entry_dir: str,
    entry_i: int,
    exit_i: int,
    entry_w: float,
    signals: pd.DataFrame,
    dates: pd.Series,
    spread_rets: pd.DataFrame,
    capital: float,
    near_expiry: pd.Series,
    exit_reason: str,
) -> dict:
    wcol = f"weight_{spread}"
    dcol = f"decision_{spread}"
    scol = f"spread_{spread}"
    ecol = f"edge_{spread}"
    idx = dates.iloc[entry_i : exit_i + 1]
    w_series = signals.loc[entry_i : exit_i, wcol].astype(float)
    if spread in spread_rets.columns:
        r = spread_rets.loc[idx, spread].reindex(idx).fillna(0.0)
        daily_pnl_pct = float((r * w_series.values).sum())
        compound = float(np.prod(1.0 + r.values * w_series.values) - 1.0)
    else:
        daily_pnl_pct = compound = 0.0

    exp_entry = near_expiry.iloc[entry_i]
    exp_exit = near_expiry.iloc[exit_i]
    same_contract = (
        pd.notna(exp_entry)
        and pd.notna(exp_exit)
        and pd.Timestamp(exp_entry).normalize() == pd.Timestamp(exp_exit).normalize()
    )

    return {
        "trade_id": trade_id,
        "spread": spread,
        "direction": entry_dir,
        "entry_date": dates.iloc[entry_i].strftime("%Y-%m-%d"),
        "exit_date": dates.iloc[exit_i].strftime("%Y-%m-%d"),
        "holding_days": exit_i - entry_i + 1,
        "entry_spread_pts": float(signals.loc[entry_i, scol])
        if scol in signals.columns and pd.notna(signals.loc[entry_i, scol])
        else np.nan,
        "exit_spread_pts": float(signals.loc[exit_i, scol])
        if scol in signals.columns and pd.notna(signals.loc[exit_i, scol])
        else np.nan,
        "entry_edge": float(signals.loc[entry_i, ecol])
        if ecol in signals.columns and pd.notna(signals.loc[entry_i, ecol])
        else np.nan,
        "exit_edge": float(signals.loc[exit_i, ecol])
        if ecol in signals.columns and pd.notna(signals.loc[exit_i, ecol])
        else np.nan,
        "entry_weight": float(entry_w),
        "avg_weight": float(w_series.mean()),
        "pnl_pct_simple": daily_pnl_pct,
        "pnl_pct_compound": compound,
        "pnl_usd": capital * daily_pnl_pct,
        "entry_decision": str(signals.loc[entry_i, dcol]),
        "exit_decision": str(signals.loc[exit_i, dcol]),
        "exit_reason": exit_reason,
        "near_expiry_entry": pd.Timestamp(exp_entry).strftime("%Y-%m-%d")
        if pd.notna(exp_entry)
        else "",
        "near_expiry_exit": pd.Timestamp(exp_exit).strftime("%Y-%m-%d")
        if pd.notna(exp_exit)
        else "",
        "same_near_contract": bool(same_contract),
        "model_type": "constant_tenor_rolled",
    }


def extract_trades(
    signals: pd.DataFrame,
    spread_rets: pd.DataFrame,
    panel: pd.DataFrame,
    *,
    capital: float,
    weight_eps: float = 0.005,
) -> pd.DataFrame:
    """
    Build trade log from daily weight / decision changes per spread.

    Splits on near-month expiry change (constant-tenor roll). Entry/exit spread
    levels are only comparable when ``same_near_contract`` is True.
    """
    signals = signals.copy()
    signals["date"] = pd.to_datetime(signals["date"])
    signals = signals.sort_values("date").reset_index(drop=True)
    dates = signals["date"]

    spread_cols = sorted(
        {c.replace("weight_", "") for c in signals.columns if c.startswith("weight_")}
    )
    rows: list[dict] = []
    trade_id = 0

    for spread in spread_cols:
        wcol = f"weight_{spread}"
        dcol = f"decision_{spread}"
        if wcol not in signals.columns:
            continue

        near_expiry = _near_expiry_series(panel, spread, dates)
        roll_today = np.zeros(len(signals), dtype=bool)
        for i in range(1, len(signals)):
            e0, e1 = near_expiry.iloc[i - 1], near_expiry.iloc[i]
            if pd.notna(e0) and pd.notna(e1) and pd.Timestamp(e0).normalize() != pd.Timestamp(e1).normalize():
                roll_today[i] = True

        in_pos = False
        entry_i: int | None = None
        entry_w = 0.0
        entry_dir = ""

        def _emit(exit_i: int, reason: str) -> None:
            nonlocal trade_id, in_pos, entry_i, entry_w, entry_dir
            if entry_i is None:
                return
            trade_id += 1
            rows.append(
                _close_trade_row(
                    trade_id=trade_id,
                    spread=spread,
                    entry_dir=entry_dir,
                    entry_i=entry_i,
                    exit_i=exit_i,
                    entry_w=entry_w,
                    signals=signals,
                    dates=dates,
                    spread_rets=spread_rets,
                    capital=capital,
                    near_expiry=near_expiry,
                    exit_reason=reason,
                )
            )
            in_pos = False
            entry_i = None

        for i in range(len(signals)):
            wt = float(signals.loc[i, wcol]) if pd.notna(signals.loc[i, wcol]) else 0.0
            dec = str(signals.loc[i, dcol]) if dcol in signals.columns else "PASS"

            if in_pos and i > 0 and roll_today[i]:
                _emit(i - 1, "near_month_roll")
                if abs(wt) >= weight_eps and dec in ("LONG", "SHORT"):
                    in_pos = True
                    entry_i = i
                    entry_w = wt
                    entry_dir = "LONG_SPREAD" if wt > 0 else "SHORT_SPREAD"
                continue

            if not in_pos:
                if abs(wt) >= weight_eps and dec in ("LONG", "SHORT"):
                    in_pos = True
                    entry_i = i
                    entry_w = wt
                    entry_dir = "LONG_SPREAD" if wt > 0 else "SHORT_SPREAD"
                continue

            sign_flip = (wt * entry_w) < 0 and abs(wt) >= weight_eps
            flat = abs(wt) < weight_eps or dec == "PASS"
            if flat or sign_flip:
                exit_i = i if flat else i - 1
                if exit_i < (entry_i or 0):
                    exit_i = entry_i or 0
                reason = "sign_flip" if sign_flip and not flat else "flat"
                _emit(exit_i, reason)
                if sign_flip and abs(wt) >= weight_eps and dec in ("LONG", "SHORT"):
                    in_pos = True
                    entry_i = i
                    entry_w = wt
                    entry_dir = "LONG_SPREAD" if wt > 0 else "SHORT_SPREAD"

        if in_pos and entry_i is not None:
            _emit(len(signals) - 1, "end_of_sample")

    if not rows:
        return pd.DataFrame()
    return pd.DataFrame(rows)


def build_equity_curve(daily: pd.DataFrame, capital: float) -> pd.DataFrame:
    d = daily.copy()
    d["date"] = pd.to_datetime(d["date"])
    d = d.sort_values("date")
    eq = d["equity_markov_usd"].astype(float)
    if eq.isna().all():
        eq = capital * (1.0 + d["daily_ret_markov"].astype(float)).cumprod()
    peak = eq.cummax()
    dd = eq / peak - 1.0
    cum_ret = eq / capital - 1.0
    out = pd.DataFrame(
        {
            "date": d["date"].dt.strftime("%Y-%m-%d"),
            "daily_ret": d["daily_ret_markov"].astype(float).values,
            "equity_usd": eq.values,
            "cumulative_return_pct": (cum_ret * 100.0).values,
            "drawdown_pct": (dd * 100.0).values,
            "gross_exposure": d.get("gross_exposure", pd.Series(np.nan, index=d.index)).values,
            "cash_weight": d.get("cash_weight", pd.Series(np.nan, index=d.index)).values,
        }
    )
    return out


def _plot_equity(curve: pd.DataFrame, png_path: Path, title: str) -> None:
    import matplotlib

    matplotlib.use("Agg")
    import matplotlib.pyplot as plt
    import matplotlib.dates as mdates

    dt = pd.to_datetime(curve["date"])
    eq = curve["equity_usd"].astype(float)
    dd = curve["drawdown_pct"].astype(float)

    fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(12, 7), sharex=True, gridspec_kw={"height_ratios": [3, 1]})
    fig.suptitle(title, fontsize=13, fontweight="bold")

    ax1.plot(dt, eq, color="#2563eb", linewidth=1.2, label="Equity ($)")
    ax1.set_ylabel("Equity ($)")
    ax1.grid(True, alpha=0.3)
    ax1.legend(loc="upper left")
    ax1.yaxis.set_major_formatter(plt.FuncFormatter(lambda x, _: f"${x:,.0f}"))

    ax2.fill_between(dt, dd, 0, color="#dc2626", alpha=0.4, label="Drawdown %")
    ax2.set_ylabel("DD (%)")
    ax2.set_xlabel("Date")
    ax2.grid(True, alpha=0.3)
    ax2.xaxis.set_major_formatter(mdates.DateFormatter("%Y"))
    ax2.xaxis.set_major_locator(mdates.YearLocator())

    plt.tight_layout()
    png_path.parent.mkdir(parents=True, exist_ok=True)
    fig.savefig(png_path, dpi=140, bbox_inches="tight")
    plt.close(fig)


def write_html_report(
    *,
    html_path: Path,
    png_path: Path,
    curve: pd.DataFrame,
    trades: pd.DataFrame,
    meta: dict,
) -> None:
    start = meta.get("start", "")
    end = meta.get("end", "")
    m = meta.get("markov_spread", {})
    cap = meta.get("capital", 100_000)

    win_trades = int((trades["pnl_usd"] > 0).sum()) if len(trades) else 0
    lose_trades = int((trades["pnl_usd"] <= 0).sum()) if len(trades) else 0
    win_rate = win_trades / len(trades) * 100 if len(trades) else 0.0

    # Show last 200 trades in HTML; full log in CSV
    trade_preview = trades.head(500) if len(trades) > 500 else trades
    trade_rows = ""
    for _, t in trade_preview.iterrows():
        pnl_cls = "pos" if t["pnl_usd"] >= 0 else "neg"
        trade_rows += (
            f"<tr>"
            f"<td>{int(t['trade_id'])}</td>"
            f"<td>{t['spread']}</td>"
            f"<td>{t['direction']}</td>"
            f"<td>{t['entry_date']}</td>"
            f"<td>{t['exit_date']}</td>"
            f"<td>{int(t['holding_days'])}</td>"
            f"<td>{t['entry_spread_pts']:.2f}</td>"
            f"<td>{t['exit_spread_pts']:.2f}</td>"
            f"<td>{t.get('near_expiry', t.get('near_expiry_entry', ''))}</td>"
            f"<td>{t.get('far_expiry', '')}</td>"
            f"<td class='{pnl_cls}'>{t['pnl_usd']:+,.0f}</td>"
            f"</tr>\n"
        )

    png_name = png_path.name
    html = f"""<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="utf-8"/>
<title>Markov VIX Spread Backtest</title>
<style>
  body {{ font-family: system-ui, sans-serif; margin: 24px; background: #0f172a; color: #e2e8f0; }}
  h1 {{ color: #f8fafc; }}
  .kpis {{ display: grid; grid-template-columns: repeat(auto-fit, minmax(160px, 1fr)); gap: 12px; margin: 20px 0; }}
  .kpi {{ background: #1e293b; padding: 14px; border-radius: 8px; }}
  .kpi label {{ font-size: 12px; color: #94a3b8; display: block; }}
  .kpi val {{ font-size: 22px; font-weight: 600; }}
  img {{ max-width: 100%; border-radius: 8px; background: #fff; }}
  table {{ width: 100%; border-collapse: collapse; font-size: 13px; margin-top: 16px; }}
  th, td {{ padding: 8px 10px; border-bottom: 1px solid #334155; text-align: left; }}
  th {{ background: #1e293b; color: #94a3b8; position: sticky; top: 0; }}
  .pos {{ color: #4ade80; }}
  .neg {{ color: #f87171; }}
  .note {{ color: #94a3b8; font-size: 13px; margin-top: 8px; }}
</style>
</head>
<body>
<h1>Markov VIX Calendar Spread Strategy</h1>
<p>{start} → {end} · ${cap:,.0f} notional · Spreads: M1−M2, M2−M3, M3−M4, M4−M5</p>

<div class="kpis">
  <div class="kpi"><label>Total return</label><val>{m.get('total_return_pct', 0):.1f}%</val></div>
  <div class="kpi"><label>CAGR</label><val>{m.get('cagr_pct', 0):.2f}%</val></div>
  <div class="kpi"><label>Sharpe</label><val>{m.get('sharpe', 0):.2f}</val></div>
  <div class="kpi"><label>Max drawdown</label><val>{m.get('max_drawdown_pct', 0):.2f}%</val></div>
  <div class="kpi"><label>Ending equity</label><val>${m.get('ending_equity_usd', cap):,.0f}</val></div>
  <div class="kpi"><label>Trades</label><val>{len(trades)}</val></div>
  <div class="kpi"><label>Win rate</label><val>{win_rate:.1f}%</val></div>
</div>

<h2>Equity curve</h2>
<img src="{png_name}" alt="Equity curve"/>

<h2>Trade log ({len(trades)} total)</h2>
<p class="note">Fixed-expiry mode locks specific VX contract pairs at entry; rolls explicitly on near-leg expiry.
Spread in/out are on the <strong>same</strong> near/far contracts for each trade row.</p>
<table>
<thead><tr>
  <th>ID</th><th>Spread</th><th>Dir</th><th>Entry</th><th>Exit</th><th>Days</th>
  <th>Spread in</th><th>Spread out</th><th>Near exp</th><th>Far exp</th><th>PnL $</th>
</tr></thead>
<tbody>
{trade_rows}
</tbody>
</table>
</body>
</html>
"""
    html_path.write_text(html)


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
    ap.add_argument("--prefix", type=Path, default=DEFAULT_PREFIX)
    ap.add_argument("--capital", type=float, default=None)
    ap.add_argument(
        "--engine-trade-log",
        action="store_true",
        help="Use trade log written by fixed-expiry engine (do not reconstruct)",
    )
    args = ap.parse_args()

    prefix = args.prefix.expanduser().resolve()
    daily_path = Path(f"{prefix}_daily.csv")
    signals_path = Path(f"{prefix}_signals.csv")
    meta_path = Path(f"{prefix}_meta.json")

    if not daily_path.is_file() or not signals_path.is_file():
        raise SystemExit(
            f"Missing {daily_path} or {signals_path}. Run run_markov_vix_spread_backtest.py first."
        )

    daily = pd.read_csv(daily_path)
    signals = pd.read_csv(signals_path)
    meta = json.loads(meta_path.read_text()) if meta_path.is_file() else {}
    capital = float(args.capital or meta.get("capital", 100_000))

    panel = pd.read_parquet(VIX_FUT_PATH)
    panel.index = pd.to_datetime(panel.index).tz_localize(None)
    master = pd.to_datetime(daily["date"])

    spread_rets = _spread_daily_returns(panel, master)

    curve = build_equity_curve(daily, capital)
    trades_path = Path(f"{prefix}_trade_log.csv")
    if args.engine_trade_log and trades_path.is_file():
        trades = pd.read_csv(trades_path)
        print(f"Using engine trade log: {trades_path} ({len(trades)} trades)")
    else:
        trades = extract_trades(signals, spread_rets, panel, capital=capital)

    eq_path = Path(f"{prefix}_equity_curve.csv")
    png_path = Path(f"{prefix}_equity_curve.png")
    html_path = Path(f"{prefix}_report.html")

    curve.to_csv(eq_path, index=False)
    trades.to_csv(trades_path, index=False)

    title = f"Markov VIX Calendar Spreads · ${capital:,.0f} · {curve['date'].iloc[0]} → {curve['date'].iloc[-1]}"
    _plot_equity(curve, png_path, title)
    write_html_report(
        html_path=html_path,
        png_path=png_path,
        curve=curve,
        trades=trades,
        meta=meta,
    )

    print(f"Equity curve: {eq_path}  ({len(curve)} days)")
    print(f"Trade log:    {trades_path}  ({len(trades)} trades)")
    print(f"Chart:        {png_path}")
    print(f"HTML report:  {html_path}")
    if len(trades):
        print(f"\nTrade summary:")
        print(f"  Wins: {int((trades['pnl_usd']>0).sum())}  Losses: {int((trades['pnl_usd']<=0).sum())}")
        print(f"  Sum PnL (attributed): ${trades['pnl_usd'].sum():,.0f}")
        print(f"  Avg hold: {trades['holding_days'].mean():.1f} days")


if __name__ == "__main__":
    main()
