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

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

import pandas as pd


CONTRACT_MULT = 100


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


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 1545 files for {root} in {theta_dir}")
    use_cols = ["quote_datetime", "expiration", "strike", "right", "bid", "ask", "delta"]
    parts: list[pd.DataFrame] = []
    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
        parts.append(sub)
    if not parts:
        return pd.DataFrame(columns=use_cols + ["session_date"])
    out = pd.concat(parts, 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")
    out["delta"] = pd.to_numeric(out["delta"], errors="coerce")
    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 _month_roll_dates_from_options(opt: pd.DataFrame) -> list[pd.Timestamp]:
    """PUTW-like systematic monthly roll: first option session date per calendar month."""
    if opt.empty:
        return []
    d = pd.to_datetime(opt["session_date"], errors="coerce").dropna().dt.normalize()
    if d.empty:
        return []
    tmp = pd.DataFrame({"d": d})
    tmp["ym"] = tmp["d"].dt.to_period("M")
    firsts = tmp.groupby("ym", as_index=False)["d"].min()["d"]
    return sorted(pd.to_datetime(firsts).dt.normalize().unique())


def run_putw_like(
    theta_dir: Path,
    root: str,
    start: pd.Timestamp,
    end: pd.Timestamp,
    *,
    dte_target: int,
    min_dte: int,
    max_dte: int,
    starting_capital: float,
) -> tuple[pd.DataFrame, float]:
    opt = _load_option_rows(theta_dir, root, start, end)
    if opt.empty:
        return pd.DataFrame(), starting_capital

    opt = opt.copy()
    opt["session_date"] = pd.to_datetime(opt["session_date"], errors="coerce").dt.normalize()
    opt = opt.loc[(opt["session_date"] >= start) & (opt["session_date"] <= end)]
    if opt.empty:
        return pd.DataFrame(), starting_capital

    by_date = {pd.Timestamp(d).normalize(): g for d, g in opt.groupby("session_date", sort=False)}
    roll_dates = _month_roll_dates_from_options(opt)

    cash = float(starting_capital)
    # Keep sizing anchored to initial capital (PUTW-style systematic notional),
    # rather than compounding contract count from prior wins/losses.
    capital_anchor = float(starting_capital)
    open_pos: dict | None = None
    trades: list[Trade] = []

    for d in roll_dates:
        d = pd.Timestamp(d).normalize()
        if open_pos is not None and d <= open_pos["expiry"]:
            continue
        chain = by_date.get(d)
        if chain is None or chain.empty:
            continue
        c = chain.copy()
        c["dte"] = (c["expiration"] - d).dt.days
        c = c.loc[(c["dte"] >= min_dte) & (c["dte"] <= max_dte)]
        if c.empty:
            continue
        exp = c.iloc[(c["dte"] - dte_target).abs().argmin()]["expiration"]
        ce = c.loc[c["expiration"] == exp].copy()
        if ce.empty:
            continue

        # PUTW-like principle: near-ATM monthly short puts.
        # Use delta closest to -0.50 when available; otherwise use middle strike.
        ce2 = ce.copy()
        if bool(ce2["delta"].notna().any()):
            ce2["delta_dist"] = (ce2["delta"] + 0.50).abs()
            pick = ce2.sort_values(["delta_dist", "strike"], ascending=[True, True]).iloc[0]
        else:
            medk = float(ce2["strike"].median())
            ce2["kdist"] = (ce2["strike"] - medk).abs()
            pick = ce2.sort_values(["kdist", "strike"], ascending=[True, True]).iloc[0]

        k = float(pick["strike"])
        exp_pick = pd.Timestamp(pick["expiration"]).normalize()
        leg = ce.loc[ce["strike"] == k]
        if leg.empty:
            continue
        prem = float(leg.iloc[0]["bid"])  # conservative execution
        if not (math.isfinite(prem) and prem > 0 and math.isfinite(k) and k > 0):
            continue

        qty = int(math.floor(capital_anchor / (k * CONTRACT_MULT)))
        if qty < 1:
            continue
        # Exit valuation from option data only: ask at expiry date, else last available ask before expiry.
        hist = opt.loc[
            (opt["expiration"] == exp_pick)
            & (opt["strike"] == k)
            & (opt["right"] == "P")
            & (opt["session_date"] <= exp_pick)
        ].sort_values("session_date")
        if hist.empty:
            continue
        close_row = hist.iloc[-1]
        close_ask = float(close_row["ask"]) if pd.notna(close_row["ask"]) else float("nan")
        if not (math.isfinite(close_ask) and close_ask > 0):
            close_mid = 0.5 * (float(close_row["bid"]) + float(close_row["ask"]))
            if math.isfinite(close_mid) and close_mid > 0:
                close_ask = close_mid
            else:
                close_ask = float(close_row["bid"]) if pd.notna(close_row["bid"]) else float("nan")
        if not (math.isfinite(close_ask) and close_ask >= 0):
            continue

        qty = int(math.floor(capital_anchor / (k * CONTRACT_MULT)))
        if qty < 1:
            continue
        pnl = (prem - close_ask) * CONTRACT_MULT * qty
        cash += pnl
        open_pos = {
            "entry_date": d,
            "expiry": exp_pick,
            "strike": float(k),
            "premium": float(prem),
            "qty": int(qty),
            "close_ask": float(close_ask),
        }
        trades.append(
            Trade(
                root=root,
                entry_date=open_pos["entry_date"],
                expiry=open_pos["expiry"],
                strike=float(open_pos["strike"]),
                spot_entry=float("nan"),
                spot_expiry=float("nan"),
                premium_per_share=float(open_pos["premium"]),
                qty=int(open_pos["qty"]),
                pnl_usd=float(pnl),
            )
        )
        open_pos = None

    tdf = pd.DataFrame([t.__dict__ for t in trades])
    return tdf, float(cash)


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


def build_equity_curve(trades_df: pd.DataFrame, start: pd.Timestamp, end: pd.Timestamp, starting_capital: float) -> pd.Series:
    """Step equity curve by realized trade PnL on expiry date."""
    idx = pd.bdate_range(start, end)
    eq = pd.Series(float(starting_capital), index=idx, dtype=float)
    if trades_df.empty:
        return eq
    t = trades_df.copy()
    t["expiry"] = pd.to_datetime(t["expiry"], errors="coerce").dt.normalize()
    pnl_by_day = t.groupby("expiry", as_index=True)["pnl_usd"].sum()
    running = float(starting_capital)
    out_vals = []
    for d in idx:
        running += float(pnl_by_day.get(pd.Timestamp(d).normalize(), 0.0))
        out_vals.append(running)
    return pd.Series(out_vals, index=idx, dtype=float)


def equal_weight_putw_portfolio_equity_series(
    theta_dir: Path,
    roots: list[str],
    start: pd.Timestamp,
    end: pd.Timestamp,
    starting_capital: float,
    *,
    dte_target: int = 30,
    min_dte: int = 20,
    max_dte: int = 45,
) -> pd.Series:
    """
    Equal-weight portfolio of :func:`run_putw_like` sleeves (same convention as ``main()``).

    Index: business days from ``start`` through ``end`` (inclusive).
    """
    start = pd.Timestamp(start).normalize()
    end = pd.Timestamp(end).normalize()
    curves: list[pd.Series] = []
    for root in roots:
        tdf, _ = run_putw_like(
            theta_dir,
            root,
            start,
            end,
            dte_target=int(dte_target),
            min_dte=int(min_dte),
            max_dte=int(max_dte),
            starting_capital=float(starting_capital),
        )
        curves.append(build_equity_curve(tdf, start, end, float(starting_capital)))
    if not curves:
        return pd.Series(dtype=float)
    curve_df = pd.concat(curves, axis=1)
    curve_df.columns = list(roots)
    return curve_df.mean(axis=1)


def main() -> None:
    ap = argparse.ArgumentParser(description="PUTW-like monthly ATM cash-secured put benchmark across tickers.")
    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("--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("--output-csv", type=Path, default=Path("RenTech/data/logs/putw_like_multi_summary.csv"))
    ap.add_argument("--trades-csv", type=Path, default=Path("RenTech/data/logs/putw_like_multi_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 = []
    trades_all = []
    per_root_trades: dict[str, pd.DataFrame] = {}
    for root in roots:
        tdf, ending_cap = run_putw_like(
            theta_dir,
            root,
            start,
            end,
            dte_target=int(args.dte_target),
            min_dte=int(args.min_dte),
            max_dte=int(args.max_dte),
            starting_capital=float(args.starting_capital),
        )
        s = summarize(tdf, start, end, float(args.starting_capital), float(ending_cap))
        rows.append({"root": root, **s})
        per_root_trades[root] = tdf
        if not tdf.empty:
            trades_all.append(tdf)

    # Equal-weight portfolio across sleeves (each sleeve starts with starting_capital;
    # portfolio is the arithmetic mean of sleeve equities).
    curves: list[pd.Series] = []
    for root in roots:
        curves.append(build_equity_curve(per_root_trades.get(root, pd.DataFrame()), start, end, float(args.starting_capital)))
    if curves:
        curve_df = pd.concat(curves, axis=1)
        curve_df.columns = roots
        port_curve = curve_df.mean(axis=1)
        port_end = float(port_curve.iloc[-1])
        years = max((end - start).days / 365.25, 1e-9)
        port_cagr = ((port_end / float(args.starting_capital)) ** (1.0 / years) - 1.0) * 100.0
        rows.append(
            {
                "root": "EQUAL_WEIGHTED",
                "trades": int(sum(len(per_root_trades.get(r, pd.DataFrame())) for r in roots)),
                "win_rate": float("nan"),
                "total_pnl_usd": float(port_end - float(args.starting_capital)),
                "ending_capital": port_end,
                "total_return_pct": (port_end / float(args.starting_capital) - 1.0) * 100.0,
                "cagr_pct": port_cagr,
                "avg_pnl_per_trade": float("nan"),
            }
        )

    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 trades_all:
        tc = pd.concat(trades_all, ignore_index=True)
        args.trades_csv.parent.mkdir(parents=True, exist_ok=True)
        tc.to_csv(args.trades_csv, index=False)
        print(f"Wrote trades:  {args.trades_csv}")


if __name__ == "__main__":
    main()

