#!/usr/bin/env python3
"""
Scan **every** JSONL trade log in a directory (or explicit glob) against **one** VRP export and
rank overlays by **low correlation** and **conditional mean** when main P&amp;L is negative.

**Overlay** rows must include ``exit_date`` and ``pnl_total`` (same as stress / straddle backtests).

**VRP** rows use ``pnl_usd`` (``vrp_backtest_theta.py --export-trades-jsonl``).

Skips files whose basename matches ``--skip-name`` (default ``vrp_trades.jsonl``).

**More overlay ideas** (implement as separate backtests, then drop JSONL into ``--overlay-dir``):

1. **Calendar / term structure** — Long near-dated, short far-dated (or reverse), same strike; rank
   legs by model ``pred`` at each tenor to express **IV curve** mispricing vs RV.
2. **Skew / risk reversal** — Long OTM put + short OTM call (or reverse) with **notional-neutral**
   delta; score with **difference** in pred between legs.
3. **Iron condor** — Short strangle + long wider wings; short **vega** + bounded risk; different
   crash profile than long vol.
4. **Put credit spread** (bullish / range) when **VIX** high — complements long-vol sleeves that
   bleed in grind-up regimes.
5. **Dispersion** — Index vs single-name variance (needs **non-SPY** option Parquet or another symbol).
6. **VVIX**-triggered sleeve — Arm overlays only when ``VVIX`` or **vol-of-vol** is elevated (use
   columns already in ``normalize_spy_df`` when available).

Example::

    python RenTech/strategy_stack/scan_iv_overlay_correlation.py \\
      --vrp-trades RenTech/data/logs/vrp_trades.jsonl \\
      --overlay-dir RenTech/data/logs \\
      --out-csv RenTech/data/logs/overlay_scan.csv
"""

from __future__ import annotations

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

import pandas as pd

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

from RenTech.strategy_stack.iv_mispricing_complement import (
    _load_jsonl,
    _pnl_series_from_trades,
    _require_jsonl,
)


def _analyze_pair(vrp_rows: list[dict], ov_rows: list[dict]) -> dict | None:
    main = _pnl_series_from_trades(vrp_rows, exit_key="exit_date", pnl_key="pnl_usd")
    ovl = _pnl_series_from_trades(ov_rows, exit_key="exit_date", pnl_key="pnl_total")
    all_idx = main.index.union(ovl.index).sort_values()
    M = main.reindex(all_idx, fill_value=0.0)
    O = ovl.reindex(all_idx, fill_value=0.0)
    mask_both = (M != 0) | (O != 0)
    if int(mask_both.sum()) < 5:
        return None
    Mc = M[mask_both]
    Oc = O[mask_both]
    corr = float(Mc.corr(Oc)) if Mc.std() > 0 and Oc.std() > 0 else float("nan")
    main_loss = M < 0
    n_loss = int(main_loss.sum())
    mean_ov_loss = float(O[main_loss].mean()) if n_loss else float("nan")
    mean_ov_ok = float(O[~main_loss].mean()) if int((~main_loss).sum()) else float("nan")
    return {
        "n_overlay_trades": len(ov_rows),
        "calendar_days_with_any_pnl": int(mask_both.sum()),
        "corr_daily_pnl_exit_day": corr,
        "days_main_pnl_negative": n_loss,
        "mean_overlay_pnl_when_main_negative": mean_ov_loss,
        "mean_overlay_pnl_when_main_non_negative": mean_ov_ok,
        "total_overlay_pnl": float(O.sum()),
    }


def main() -> None:
    ap = argparse.ArgumentParser(description="Batch-scan overlay JSONLs vs VRP for correlation")
    ap.add_argument("--vrp-trades", type=Path, required=True)
    ap.add_argument(
        "--overlay-dir",
        type=Path,
        default=None,
        help="Directory containing *.jsonl overlay trade logs",
    )
    ap.add_argument(
        "--overlay-glob",
        type=str,
        default="*.jsonl",
        help="Glob under --overlay-dir (default all jsonl)",
    )
    ap.add_argument(
        "--skip-name",
        action="append",
        default=["vrp_trades.jsonl"],
        help="Basename(s) to skip (repeatable)",
    )
    ap.add_argument("--out-csv", type=Path, default=None)
    ap.add_argument("--min-overlay-trades", type=int, default=3)
    args = ap.parse_args()

    vrp_path = _require_jsonl(args.vrp_trades, hint="VRP --export-trades-jsonl")
    vrp_rows = _load_jsonl(vrp_path)

    paths: list[Path] = []
    if args.overlay_dir is not None:
        od = args.overlay_dir.expanduser().resolve()
        if not od.is_dir():
            print(f"ERROR: not a directory: {od}", file=sys.stderr)
            sys.exit(1)
        paths = sorted(od.glob(args.overlay_glob))
    else:
        print("ERROR: pass --overlay-dir", file=sys.stderr)
        sys.exit(1)

    skip = {s.lower() for s in args.skip_name}
    rows_out: list[dict] = []

    for p in paths:
        if not p.is_file():
            continue
        if p.name.lower() in skip:
            continue
        try:
            ov_rows = _load_jsonl(p)
        except (OSError, json.JSONDecodeError) as e:
            rows_out.append({"file": str(p), "error": str(e)})
            continue
        if len(ov_rows) < int(args.min_overlay_trades):
            rows_out.append({"file": str(p), "error": "too_few_trades"})
            continue
        r = _analyze_pair(vrp_rows, ov_rows)
        if r is None:
            rows_out.append({"file": str(p), "error": "too_few_overlapping_days"})
            continue
        r["file"] = str(p)
        r["basename"] = p.name
        rows_out.append(r)

    # Sort by abs correlation ascending (more uncorrelated first), then by mean when main negative
    def sort_key(x: dict) -> tuple:
        if "corr_daily_pnl_exit_day" not in x:
            return (9e9, 0.0)
        c = x["corr_daily_pnl_exit_day"]
        cc = abs(c) if isinstance(c, float) and math.isfinite(c) else 9e9
        m = x.get("mean_overlay_pnl_when_main_negative")
        mm = m if isinstance(m, float) and math.isfinite(m) else 0.0
        return (cc, -mm)

    ranked = sorted([x for x in rows_out if "corr_daily_pnl_exit_day" in x], key=sort_key)
    tail = [x for x in rows_out if "corr_daily_pnl_exit_day" not in x]

    print(json.dumps({"n_scanned": len(paths), "n_ok": len(ranked), "n_errors": len(tail)}, indent=2))
    for r in ranked[:25]:
        print(
            f"{r.get('basename',''):40}  corr={r.get('corr_daily_pnl_exit_day', float('nan')):+.4f}  "
            f"ov|main<0={r.get('mean_overlay_pnl_when_main_negative', float('nan')):+.2f}  "
            f"total_ov={r.get('total_overlay_pnl', 0):+.0f}"
        )

    if args.out_csv:
        outp = args.out_csv.expanduser()
        outp.parent.mkdir(parents=True, exist_ok=True)
        all_rows = ranked + tail
        if all_rows:
            keys = set()
            for r in all_rows:
                keys.update(r.keys())
            fieldnames = sorted(keys)
            with outp.open("w", newline="", encoding="utf-8") as f:
                w = csv.DictWriter(f, fieldnames=fieldnames, extrasaction="ignore")
                w.writeheader()
                for r in all_rows:
                    w.writerow(r)
            print(f"Wrote {outp}")


if __name__ == "__main__":
    main()
