#!/usr/bin/env python3
"""
Train an **XGBoost** model on Theta Parquet rows to estimate **forward realized vol minus IV**
(a simple **vol risk premium** proxy at the contract level).

**Label (regression):** ``rv_fwd_H - implied_vol`` where ``rv_fwd_H`` is annualized realized vol
of SPY over the next ``H`` **trading** days (log-return std × sqrt(252)). Positive label ⇒
forward realized vol **higher** than option-implied vol (IV looked **cheap** ex post).

**Features:** ``dte``, ``moneyness`` (K/S), ``abs_log_moneyness``, ``is_call``, ``implied_vol``,
``iv_inferred`` (1 if IV was back-solved in enrich step), ``vix_close`` (by session date).

This does **not** place trades. Use predictions as a **research signal**: e.g. rank contracts by
predicted edge, then build **vega-neutral** or **delta-neutral** structures (straddle / strangle /
ratio spreads) in a separate execution layer. True delta neutrality requires live greeks and
position sizing.

Dependencies: pandas, numpy, pyarrow, xgboost, scikit-learn, joblib, yfinance.

Usage::

    python RenTech/strategy_stack/train_vol_mispricing_xgb.py --theta-dir RenTech/data/theta_chunks --horizon 5
    python RenTech/strategy_stack/train_vol_mispricing_xgb.py --max-rows 200000 --out RenTech/data/models/vol_mispricing_xgb.json
    python RenTech/strategy_stack/train_vol_mispricing_xgb.py --train-end 2022-12-31   # time split (less lookahead bias)

Prefers ``spy_1545_*_ivfilled.parquet`` over the same month without suffix when both exist.
"""

from __future__ import annotations

import argparse
import json
import math
import sys
from pathlib import Path

_REPO = Path(__file__).resolve().parents[2]
if str(_REPO) not in sys.path:
    sys.path.insert(0, str(_REPO))

import numpy as np
import pandas as pd

from RenTech.core.theta_chunks_loader import _session_dates_series
from RenTech.strategy_stack.vrp_backtester import load_spy_vix_from_yfinance, normalize_spy_df

FEATURE_COLUMNS = [
    "dte",
    "moneyness",
    "abs_log_moneyness",
    "is_call",
    "implied_vol",
    "iv_inferred",
    "vix_close",
]
TARGET_COLUMN = "vrp_edge"  # rv_fwd - iv
# Not a model feature; used for --train-end filtering only.
_META_SESSION = "_session"


def resolve_month_parquet(theta_dir: Path | str, ts: pd.Timestamp) -> Path | None:
    """
    Single monthly chunk path for ``ts``, preferring ``*_ivfilled.parquet`` when present.
    """
    theta_dir = Path(theta_dir).expanduser()
    y, m = int(ts.year), int(ts.month)
    iv = theta_dir / f"spy_1545_{y:04d}_{m:02d}_ivfilled.parquet"
    base = theta_dir / f"spy_1545_{y:04d}_{m:02d}.parquet"
    if iv.is_file():
        return iv
    if base.is_file():
        return base
    return None


def resolve_theta_parquet_paths(theta_dir: Path) -> list[Path]:
    """
    One file per calendar month. If ``spy_1545_YYYY_MM_ivfilled.parquet`` exists, use it;
    otherwise ``spy_1545_YYYY_MM.parquet``. Avoids double-counting when both are present.
    """
    theta_dir = theta_dir.expanduser()
    stems = set()
    for p in theta_dir.glob("spy_1545_*.parquet"):
        s = p.stem
        if s.endswith("_ivfilled"):
            stems.add(s[: -len("_ivfilled")])
        else:
            stems.add(s)
    out: list[Path] = []
    for stem in sorted(stems):
        iv = theta_dir / f"{stem}_ivfilled.parquet"
        base = theta_dir / f"{stem}.parquet"
        if iv.is_file():
            out.append(iv)
        elif base.is_file():
            out.append(base)
    return out


def _scale_strike(strike: float, spy_px: float) -> float:
    k = float(strike)
    if k < 150:
        k *= 10.0
    return k


def _featurize_row(
    r: pd.Series,
    sn: pd.Timestamp,
    spy_df: pd.DataFrame,
    date_to_i: dict[pd.Timestamp, int],
) -> dict[str, float] | None:
    """
    Build one training/inference row (feature dict) or None if filters fail.
    ``sn`` must be normalized session date present in ``spy_df``.
    """
    if sn not in date_to_i:
        return None
    iv = float(pd.to_numeric(r.get("implied_vol"), errors="coerce"))
    if not (math.isfinite(iv) and iv > 0):
        return None
    bid = float(r["bid"]) if pd.notna(r.get("bid")) else float("nan")
    ask = float(r["ask"]) if pd.notna(r.get("ask")) else float("nan")
    if not (math.isfinite(bid) and math.isfinite(ask) and bid > 0 and ask > 0):
        return None

    spy_px = float(spy_df.loc[sn, "close"])
    K = _scale_strike(float(r["strike"]), spy_px)
    exp = pd.Timestamp(r["expiration"]).normalize()
    dte = int((exp - sn).days)
    if dte < 1:
        return None
    m = K / spy_px
    opt = str(r.get("right", "")).strip().upper()
    is_call = 1.0 if opt.startswith("C") else 0.0
    vix = float(spy_df.loc[sn, "vix_close"])
    iv_raw = r.get("iv_inferred", False)
    iv_inf = 1.0 if (bool(iv_raw) or str(iv_raw).lower() in ("1", "true")) else 0.0

    return {
        "dte": float(dte),
        "moneyness": float(m),
        "abs_log_moneyness": float(abs(math.log(m))),
        "is_call": is_call,
        "implied_vol": iv,
        "iv_inferred": iv_inf,
        "vix_close": vix,
    }


def forward_rv_ann(spy_close: np.ndarray, start_idx: int, horizon: int) -> float:
    """Annualized RV from `horizon` daily log returns starting at SPY move after start_idx."""
    if start_idx + horizon >= len(spy_close):
        return float("nan")
    s = spy_close[start_idx : start_idx + horizon + 1]
    if np.any(s <= 0):
        return float("nan")
    rets = np.log(s[1:] / s[:-1])
    if len(rets) < horizon:
        return float("nan")
    return float(np.sqrt(252.0) * np.std(rets, ddof=1))


def build_frame_from_parquets(
    paths: list[Path],
    spy_df: pd.DataFrame,
    *,
    horizon: int,
    max_rows: int,
    rng: np.random.Generator,
    per_file_cap: int | None = None,
) -> pd.DataFrame:
    """Sample rows from Parquet files; compute label where forward window exists."""
    idx_dates = spy_df.index
    close_s = spy_df["close"].astype(float).values
    date_to_i = {pd.Timestamp(ix).normalize(): j for j, ix in enumerate(idx_dates)}

    rows: list[dict] = []
    n_paths = max(1, len(paths))
    auto_cap = (
        max(400, int(max_rows) // max(6, min(n_paths, 40)))
        if max_rows > 0
        else 10**15
    )
    cap = per_file_cap if per_file_cap is not None else auto_cap

    for p in paths:
        if max_rows > 0 and len(rows) >= max_rows:
            break
        need = max_rows - len(rows) if max_rows > 0 else 10**15
        df = pd.read_parquet(p)
        if df.empty:
            continue
        if max_rows > 0 and len(df) > min(need, cap):
            n_take = min(len(df), max(need * 2, min(cap, 50_000)), cap)
            df = df.sample(n=n_take, random_state=int(rng.integers(1 << 31)))
        qd = _session_dates_series(df["quote_datetime"])
        sess = pd.to_datetime(qd).dt.normalize()

        for k in range(len(df)):
            if max_rows > 0 and len(rows) >= max_rows:
                break
            r = df.iloc[k]
            sn = sess.iloc[k]
            if pd.isna(sn):
                continue
            sn = pd.Timestamp(sn).normalize()
            if sn not in date_to_i:
                continue
            i0 = date_to_i[sn]
            feats = _featurize_row(r, sn, spy_df, date_to_i)
            if feats is None:
                continue

            rv = forward_rv_ann(close_s, i0, horizon)
            if not math.isfinite(rv):
                continue
            iv = feats["implied_vol"]

            rows.append(
                {
                    _META_SESSION: sn,
                    **feats,
                    TARGET_COLUMN: float(rv - iv),
                }
            )

    return pd.DataFrame(rows)


def featurize_dataframe(df: pd.DataFrame, spy_df: pd.DataFrame) -> pd.DataFrame:
    """
    Vector-free path for inference: same filters as training (no label). Each row includes
    ``_row`` = integer position in ``df`` for aligning predictions.
    """
    idx_dates = spy_df.index
    date_to_i = {pd.Timestamp(ix).normalize(): j for j, ix in enumerate(idx_dates)}
    qd = _session_dates_series(df["quote_datetime"])
    sess = pd.to_datetime(qd).dt.normalize()
    rows: list[dict] = []
    for k in range(len(df)):
        r = df.iloc[k]
        sn = sess.iloc[k]
        if pd.isna(sn):
            continue
        sn = pd.Timestamp(sn).normalize()
        feats = _featurize_row(r, sn, spy_df, date_to_i)
        if feats is None:
            continue
        rows.append({**feats, "_row": k})
    return pd.DataFrame(rows)


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--theta-dir", type=Path, default=_REPO / "RenTech" / "data" / "theta_chunks")
    ap.add_argument("--horizon", type=int, default=5, help="Forward trading days for realized vol")
    ap.add_argument("--max-rows", type=int, default=300_000, help="Cap training rows (0 = no cap)")
    ap.add_argument("--test-frac", type=float, default=0.15)
    ap.add_argument("--out", type=Path, default=_REPO / "RenTech" / "data" / "models" / "vol_mispricing_xgb.json")
    ap.add_argument("--seed", type=int, default=42)
    ap.add_argument(
        "--train-end",
        type=str,
        default="",
        help="Optional YYYY-MM-DD: train on rows with session date <= this, test on rows after (time split).",
    )
    args = ap.parse_args()

    try:
        import joblib
        from sklearn.metrics import mean_absolute_error, r2_score
        from sklearn.model_selection import train_test_split
        import xgboost as xgb
    except ImportError as e:
        print("Install: pip install xgboost scikit-learn joblib", file=sys.stderr)
        raise

    theta_dir = args.theta_dir.expanduser()
    paths = resolve_theta_parquet_paths(theta_dir)
    if not paths:
        print(f"No spy_1545_*.parquet under {theta_dir}", file=sys.stderr)
        sys.exit(1)

    # yfinance panel covering chunk dates
    ts0 = pd.read_parquet(paths[0], columns=["quote_datetime"])
    q = _session_dates_series(ts0["quote_datetime"]).min()
    ts1 = pd.read_parquet(paths[-1], columns=["quote_datetime"])
    q1 = _session_dates_series(ts1["quote_datetime"]).max()
    yf_start = (pd.Timestamp(q) - pd.Timedelta(days=400)).strftime("%Y-%m-%d")
    yf_end = (pd.Timestamp(q1) + pd.Timedelta(days=60)).strftime("%Y-%m-%d")
    spy_df = normalize_spy_df(load_spy_vix_from_yfinance(yf_start, yf_end))

    rng = np.random.default_rng(args.seed)
    max_rows = int(args.max_rows) if args.max_rows > 0 else 10**9
    train_end_s = (args.train_end or "").strip()
    paths_work = list(paths)
    if train_end_s:
        # Otherwise early months dominate the sample and time-split test sets can be empty.
        rng.shuffle(paths_work)
    frame = build_frame_from_parquets(
        paths_work, spy_df, horizon=int(args.horizon), max_rows=max_rows, rng=rng
    )
    if frame.empty or len(frame) < 500:
        print(f"Too few rows after filters: {len(frame)}", file=sys.stderr)
        sys.exit(1)

    X = frame[FEATURE_COLUMNS].astype(np.float64)
    y = frame[TARGET_COLUMN].astype(np.float64)

    if train_end_s:
        tcut = pd.Timestamp(train_end_s).normalize()
        sess = pd.to_datetime(frame[_META_SESSION]).dt.normalize()
        train_mask = sess <= tcut
        n_tr, n_te = int(train_mask.sum()), int((~train_mask).sum())
        min_te = 50 if max_rows > 0 and max_rows < 10_000 else 100
        if n_tr < 400 or n_te < min_te:
            print(
                f"Time split: need >=400 train and >={min_te} test rows; got "
                f"{n_tr} train, {n_te} test. Try larger --max-rows or shuffle-friendly --train-end.",
                file=sys.stderr,
            )
            sys.exit(1)
        X_train = X.loc[train_mask]
        X_test = X.loc[~train_mask]
        y_train = y.loc[train_mask]
        y_test = y.loc[~train_mask]
    else:
        X_train, X_test, y_train, y_test = train_test_split(
            X, y, test_size=float(args.test_frac), random_state=args.seed
        )
    model = xgb.XGBRegressor(
        n_estimators=300,
        max_depth=6,
        learning_rate=0.05,
        subsample=0.8,
        colsample_bytree=0.85,
        n_jobs=-1,
        random_state=args.seed,
    )
    model.fit(X_train, y_train)
    pred = model.predict(X_test)
    mae = mean_absolute_error(y_test, pred)
    r2 = r2_score(y_test, pred)
    importance = dict(zip(FEATURE_COLUMNS, model.feature_importances_.tolist(), strict=True))

    args.out.parent.mkdir(parents=True, exist_ok=True)
    joblib.dump(
        {
            "model": model,
            "feature_columns": FEATURE_COLUMNS,
            "target": TARGET_COLUMN,
            "horizon": int(args.horizon),
            "mae": mae,
            "r2": r2,
            "feature_importance": importance,
            "train_end": train_end_s or None,
            "split": "time" if train_end_s else "random",
        },
        args.out.with_suffix(".joblib"),
    )
    meta = {
        "feature_columns": FEATURE_COLUMNS,
        "target": TARGET_COLUMN,
        "horizon": int(args.horizon),
        "mae": mae,
        "r2": r2,
        "n_rows": len(frame),
        "artifact": str(args.out.with_suffix(".joblib")),
        "feature_importance": importance,
        "train_end": train_end_s or None,
        "split": "time" if train_end_s else "random",
    }
    args.out.write_text(json.dumps(meta, indent=2), encoding="utf-8")

    print(json.dumps(meta, indent=2))


if __name__ == "__main__":
    main()
