#!/usr/bin/env python3
from __future__ import annotations

import argparse
import math
from dataclasses import dataclass
from pathlib import Path

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


@dataclass
class ClosedTrade:
    root: str
    entry_date: pd.Timestamp
    expiry: pd.Timestamp
    strike: float
    spot_entry: float
    spot_expiry: float
    premium_per_share: float
    pnl_usd: float


def _ticker_for_root(root: str) -> str:
    ru = root.upper()
    if ru == "SPX":
        return "^SPX"
    if ru == "VIX":
        return "^VIX"
    return ru


def _session_dates_series(quote_datetime: pd.Series) -> pd.Series:
    qd = pd.to_datetime(quote_datetime, utc=False)
    if qd.dt.tz is None:
        qd = qd.dt.tz_localize("America/New_York", ambiguous="NaT", nonexistent="shift_forward")
    else:
        qd = qd.dt.tz_convert("America/New_York")
    return pd.to_datetime(qd.dt.date)


def _load_option_rows(theta_dir: Path, root: str, start: pd.Timestamp, end: pd.Timestamp) -> pd.DataFrame:
    files = sorted(theta_dir.glob(f"{root.lower()}_1545_*.parquet"))
    if not files:
        raise FileNotFoundError(f"No files for {root} in {theta_dir}")
    dfs: list[pd.DataFrame] = []
    use_cols = ["quote_datetime", "expiration", "strike", "right", "bid", "ask"]
    for p in files:
        df = pd.read_parquet(p, columns=use_cols)
        if df.empty:
            continue
        sd = _session_dates_series(df["quote_datetime"])
        m = (sd >= start) & (sd <= end)
        if not bool(m.any()):
            continue
        sub = df.loc[m].copy()
        sub["session_date"] = sd.loc[m].values
        dfs.append(sub)
    if not dfs:
        return pd.DataFrame(columns=use_cols + ["session_date"])
    out = pd.concat(dfs, ignore_index=True)
    out["expiration"] = pd.to_datetime(out["expiration"], errors="coerce").dt.normalize()
    out["right"] = out["right"].astype(str).str.upper().str.strip().str[0]
    out["bid"] = pd.to_numeric(out["bid"], errors="coerce")
    out["ask"] = pd.to_numeric(out["ask"], errors="coerce")
    out["strike"] = pd.to_numeric(out["strike"], errors="coerce")
    # Theta strike normalization used elsewhere in repo.
    if bool(out["strike"].max(skipna=True) < 150):
        out["strike"] = out["strike"] * 10.0
    out = out.dropna(subset=["session_date", "expiration", "strike", "bid", "ask"])
    out = out.loc[(out["right"] == "P") & (out["bid"] > 0) & (out["ask"] > 0)]
    return out


def _load_price_panel(root: str, start: pd.Timestamp, end: pd.Timestamp) -> pd.DataFrame:
    t = _ticker_for_root(root)
    h = yf.download(
        t,
        start=(start - pd.Timedelta(days=400)).strftime("%Y-%m-%d"),
        end=(end + pd.Timedelta(days=5)).strftime("%Y-%m-%d"),
        auto_adjust=True,
        progress=False,
        interval="1d",
    )
    if h is None or h.empty:
        raise RuntimeError(f"No yfinance history for {root}/{t}")
    c = h["Close"]
    if isinstance(c, pd.DataFrame):
        c = c.iloc[:, 0]
    close = pd.to_numeric(c, errors="coerce").dropna().rename("close")
    px = close.to_frame()
    px.index = pd.to_datetime(px.index).normalize()
    px["sma200"] = px["close"].rolling(200, min_periods=200).mean()
    return px


def run_root(
    theta_dir: Path,
    root: str,
    start: pd.Timestamp,
    end: pd.Timestamp,
    *,
    otm_pct: float,
    dte_target: int,
    min_dte: int,
    max_dte: int,
    require_sma200: bool,
) -> tuple[pd.DataFrame, list[ClosedTrade]]:
    opt = _load_option_rows(theta_dir, root, start, end)
    px = _load_price_panel(root, start, end)
    px = px.loc[(px.index >= start) & (px.index <= end)].copy()
    if px.empty:
        return pd.DataFrame(), []

    by_date = {d: g for d, g in opt.groupby("session_date", sort=False)}
    open_pos: dict | None = None
    trades: list[ClosedTrade] = []
    last_entry_month: tuple[int, int] | None = None

    for d, row in px.iterrows():
        spot = float(row["close"])
        sma = float(row["sma200"]) if pd.notna(row["sma200"]) else float("nan")

        if open_pos is not None and d >= open_pos["expiry"]:
            exp = pd.Timestamp(open_pos["expiry"]).normalize()
            exp_spot = float(px.loc[px.index[px.index <= exp][-1], "close"]) if bool((px.index <= exp).any()) else spot
            intrinsic = max(open_pos["strike"] - exp_spot, 0.0)
            pnl = (open_pos["premium"] - intrinsic) * 100.0
            trades.append(
                ClosedTrade(
                    root=root,
                    entry_date=open_pos["entry_date"],
                    expiry=exp,
                    strike=float(open_pos["strike"]),
                    spot_entry=float(open_pos["spot_entry"]),
                    spot_expiry=float(exp_spot),
                    premium_per_share=float(open_pos["premium"]),
                    pnl_usd=float(pnl),
                )
            )
            open_pos = None

        if open_pos is not None:
            continue
        if require_sma200 and not (math.isfinite(sma) and spot > sma):
            continue

        ym = (d.year, d.month)
        if last_entry_month == ym:
            continue
        daily = by_date.get(pd.Timestamp(d).normalize())
        if daily is None or daily.empty:
            continue

        cands = daily.copy()
        cands["dte"] = (cands["expiration"] - pd.Timestamp(d).normalize()).dt.days
        cands = cands.loc[(cands["dte"] >= min_dte) & (cands["dte"] <= max_dte)]
        if cands.empty:
            continue
        exp = cands.iloc[(cands["dte"] - dte_target).abs().argmin()]["expiration"]
        chain = cands.loc[cands["expiration"] == exp].copy()
        if chain.empty:
            continue
        target_k = spot * (1.0 - otm_pct)
        below = chain.loc[chain["strike"] <= target_k]
        if below.empty:
            k = float(chain.iloc[(chain["strike"] - target_k).abs().argmin()]["strike"])
        else:
            k = float(below["strike"].max())
        leg = chain.loc[chain["strike"] == k]
        if leg.empty:
            continue
        premium = float(leg.iloc[0]["bid"])  # conservative sell at bid
        if not (math.isfinite(premium) and premium > 0):
            continue

        open_pos = {
            "entry_date": pd.Timestamp(d).normalize(),
            "expiry": pd.Timestamp(exp).normalize(),
            "strike": float(k),
            "premium": float(premium),
            "spot_entry": float(spot),
        }
        last_entry_month = ym

    if open_pos is not None:
        exp = pd.Timestamp(open_pos["expiry"]).normalize()
        exp_spot = float(px.iloc[-1]["close"])
        intrinsic = max(open_pos["strike"] - exp_spot, 0.0)
        pnl = (open_pos["premium"] - intrinsic) * 100.0
        trades.append(
            ClosedTrade(
                root=root,
                entry_date=open_pos["entry_date"],
                expiry=exp,
                strike=float(open_pos["strike"]),
                spot_entry=float(open_pos["spot_entry"]),
                spot_expiry=float(exp_spot),
                premium_per_share=float(open_pos["premium"]),
                pnl_usd=float(pnl),
            )
        )

    tdf = pd.DataFrame([t.__dict__ for t in trades])
    if tdf.empty:
        return tdf, trades
    tdf["entry_date"] = pd.to_datetime(tdf["entry_date"]).dt.normalize()
    tdf["expiry"] = pd.to_datetime(tdf["expiry"]).dt.normalize()
    return tdf, trades


def summarize(trades_df: pd.DataFrame, start: pd.Timestamp, end: pd.Timestamp, starting_capital: float) -> dict:
    if trades_df.empty:
        years = max((end - start).days / 365.25, 1e-9)
        return {
            "trades": 0,
            "win_rate": 0.0,
            "total_pnl_usd": 0.0,
            "ending_capital": starting_capital,
            "total_return_pct": 0.0,
            "cagr_pct": 0.0 if years <= 0 else 0.0,
            "avg_pnl_per_trade": 0.0,
        }
    total_pnl = float(trades_df["pnl_usd"].sum())
    wins = float((trades_df["pnl_usd"] > 0).mean())
    end_cap = float(starting_capital + total_pnl)
    years = max((end - start).days / 365.25, 1e-9)
    cagr = (end_cap / starting_capital) ** (1.0 / years) - 1.0 if end_cap > 0 else -1.0
    return {
        "trades": int(len(trades_df)),
        "win_rate": wins,
        "total_pnl_usd": total_pnl,
        "ending_capital": end_cap,
        "total_return_pct": (end_cap / starting_capital - 1.0) * 100.0,
        "cagr_pct": cagr * 100.0,
        "avg_pnl_per_trade": float(trades_df["pnl_usd"].mean()),
    }


def main() -> None:
    ap = argparse.ArgumentParser(description="Benchmark monthly 2% OTM short puts with close>SMA200 filter.")
    ap.add_argument("--theta-dir", type=Path, default=Path("RenTech/data/theta_chunks"))
    ap.add_argument("--roots", type=str, default="SPY,TLT,GLD,IWM,QQQ,USO")
    ap.add_argument("--start-date", type=str, default="2016-04-01")
    ap.add_argument("--end-date", type=str, default="2026-04-30")
    ap.add_argument("--otm-pct", type=float, default=0.02)
    ap.add_argument("--dte-target", type=int, default=30)
    ap.add_argument("--min-dte", type=int, default=20)
    ap.add_argument("--max-dte", type=int, default=45)
    ap.add_argument("--starting-capital", type=float, default=100000.0)
    ap.add_argument(
        "--no-sma-gate",
        action="store_true",
        help="Disable the close>SMA(200) entry filter.",
    )
    ap.add_argument("--output-csv", type=Path, default=Path("RenTech/data/logs/monthly_putwrite_sma200_summary.csv"))
    ap.add_argument("--trades-csv", type=Path, default=Path("RenTech/data/logs/monthly_putwrite_sma200_trades.csv"))
    args = ap.parse_args()

    roots = [r.strip().upper() for r in args.roots.split(",") if r.strip()]
    start = pd.Timestamp(args.start_date).normalize()
    end = pd.Timestamp(args.end_date).normalize()
    theta_dir = args.theta_dir.expanduser().resolve()

    rows: list[dict] = []
    all_trades: list[pd.DataFrame] = []
    for root in roots:
        tdf, _ = run_root(
            theta_dir,
            root,
            start,
            end,
            otm_pct=float(args.otm_pct),
            dte_target=int(args.dte_target),
            min_dte=int(args.min_dte),
            max_dte=int(args.max_dte),
            require_sma200=not bool(args.no_sma_gate),
        )
        s = summarize(tdf, start, end, float(args.starting_capital))
        rows.append({"root": root, **s})
        if not tdf.empty:
            all_trades.append(tdf)

    out = pd.DataFrame(rows).sort_values("cagr_pct", ascending=False)
    print(out.to_string(index=False, float_format=lambda x: f"{x:,.4f}"))

    args.output_csv.parent.mkdir(parents=True, exist_ok=True)
    out.to_csv(args.output_csv, index=False)
    print(f"\nWrote summary: {args.output_csv}")

    if all_trades:
        tcat = pd.concat(all_trades, ignore_index=True)
        args.trades_csv.parent.mkdir(parents=True, exist_ok=True)
        tcat.to_csv(args.trades_csv, index=False)
        print(f"Wrote trades:  {args.trades_csv}")


if __name__ == "__main__":
    main()

