#!/usr/bin/env python3
"""
Full **~10y** trade package for the **Sharpe-optimized** merged book (VRP + IV + VXX).

1. **Re-runs** the VRP Theta engine with the same preset as ``run_vrp_low_dd_vxx_bundle.py``
   (low-DD overlap: vol scaling off, R2 crossover on, DD scaling on, overlap on).
2. Loads **existing** IV / VXX engine JSONLs (no Theta re-run for overlays) and applies the **same
   linear PnL scale** the merge uses: ``pnl_scaled = pnl_raw * (capital_sleeve / sleeve_risk_ref)``.
3. Writes per-sleeve CSVs, one **combined** CSV sorted by exit date, VRP JSONL, merged equity CSV,
   and a **manifest** with scale factors and reproduction commands.

Default sleeve budgets match the 2026-05-08 run: ``--sharpe-max-dd-pct 10`` interior optimum
(put ~0.523%%, straddle ~5.75%%, RR ~3.39%%, VXX ~23.2%% of ``--capital``). Override with
``--capital-*-pct`` flags.

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 .venv/bin/python \\
      RenTech/strategy_stack/export_sharpe_opt_fullbook_trades.py \\
      --overlap-slice-contracts 2 --no-progress
"""
from __future__ import annotations

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

_REPO_ROOT = Path(__file__).resolve().parents[2]
_LOGS = _REPO_ROOT / "RenTech" / "data" / "logs"
_THETA_DIR = _REPO_ROOT / "RenTech" / "data" / "theta_chunks"

if str(_REPO_ROOT) not in sys.path:
    sys.path.insert(0, str(_REPO_ROOT))

import pandas as pd

from RenTech.core.theta_chunks_loader import ThetaChunksLoader, theta_chunks_date_bounds
from RenTech.strategy_stack.export_theta_full_trade_log import _write_csv as _write_vrp_csv
from RenTech.strategy_stack.iv_mispricing_complement import _load_jsonl
from RenTech.strategy_stack.portfolio_vrp_plus_vxx import (
    _infer_broker_risk_overlay,
    _infer_broker_risk_vxx,
    _metrics_block,
    _overlay_sleeve_risk_reference_usd,
    execute_portfolio_merge,
)
from RenTech.strategy_stack.run_vrp_low_dd_vxx_bundle import _write_vrp_jsonl
from RenTech.strategy_stack.vrp_backtester import VRPBacktester, load_spy_vix_from_yfinance, normalize_spy_df, trading_days_intersecting_spy
from RenTech.strategy_stack.vrp_strategy_config import (
    DEFAULT_STRATEGY_CONFIG_PATH,
    apply_strategy_params_to_vrp_backtester_module,
    load_strategy_config_file,
)

# Interior Sharpe optimum under --sharpe-max-dd-pct 10 (full VRP scale), 2026-05-08.
_DEFAULT_PUT_FRAC = 0.00523407
_DEFAULT_STRADDLE_FRAC = 0.05749554
_DEFAULT_RR_FRAC = 0.03388598
_DEFAULT_VXX_FRAC = 0.23201529


def _scale_rows(
    rows: list[dict],
    *,
    sleeve: str,
    pnl_key: str,
    capital_usd: float,
    ref_usd: float,
    exit_key: str = "exit_date",
) -> list[dict]:
    ref = max(float(ref_usd), 1e-12)
    s = float(capital_usd) / ref
    out: list[dict] = []
    for r in rows:
        raw = float(r.get(pnl_key, 0.0) or 0.0)
        ex = r.get(exit_key)
        en = r.get("entry_date", "")
        out.append(
            {
                "sleeve": sleeve,
                "entry_date": str(en)[:10] if en is not None else "",
                "exit_date": str(ex)[:10] if ex is not None else "",
                "pnl_usd_scaled": round(raw * s, 6),
                "pnl_usd_raw_in_log": round(raw, 6),
                "sleeve_scale_factor": round(s, 10),
                "sleeve_risk_ref_usd": round(ref, 2),
                "sleeve_capital_usd": round(float(capital_usd), 2),
            }
        )
    return out


def main() -> None:
    ap = argparse.ArgumentParser(description="Export full-book trade logs for Sharpe-opt merge preset.")
    ap.add_argument("--theta-dir", type=Path, default=_THETA_DIR)
    ap.add_argument("--capital", type=float, default=100_000.0, help="Account / merge total capital.")
    ap.add_argument("--vrp-ref", type=float, default=100_000.0)
    ap.add_argument("--capital-vrp-pct", type=float, default=1.0, metavar="FRAC", help="VRP PnL numerator = FRAC × --capital (default 1).")
    ap.add_argument("--capital-put-pct", type=float, default=_DEFAULT_PUT_FRAC)
    ap.add_argument("--capital-straddle-pct", type=float, default=_DEFAULT_STRADDLE_FRAC)
    ap.add_argument("--capital-risk-reversal-pct", type=float, default=_DEFAULT_RR_FRAC)
    ap.add_argument("--capital-vxx-pct", type=float, default=_DEFAULT_VXX_FRAC)
    ap.add_argument("--vxx-bear-pct", type=float, default=90.0)
    ap.add_argument("--vxx-call-pct", type=float, default=10.0)
    ap.add_argument("--overlap-slice-contracts", type=int, default=2, metavar="N")
    ap.add_argument("--start", type=str, default="", help="Optional YYYY-MM-DD floor.")
    ap.add_argument("--end", type=str, default="", help="Optional YYYY-MM-DD cap.")
    ap.add_argument("--max-days", type=int, default=0, help="If >0, smoke test first N sessions.")
    ap.add_argument("--no-progress", action="store_true")
    ap.add_argument(
        "--out-dir",
        type=Path,
        default=_LOGS,
        help="Directory for all outputs (default: RenTech/data/logs/).",
    )
    ap.add_argument(
        "--out-prefix",
        type=str,
        default="sharpe_opt_10dd_fullbook_ov2",
        help="Filename prefix for outputs.",
    )
    ap.add_argument("--put-trades", type=Path, default=_LOGS / "engine_iv_otm_put_ivx1.jsonl")
    ap.add_argument("--straddle-trades", type=Path, default=_LOGS / "engine_iv_straddle_ivx1.jsonl")
    ap.add_argument("--risk-reversal-trades", type=Path, default=_LOGS / "engine_iv_risk_reversal_ivx1.jsonl")
    ap.add_argument(
        "--vxx-bear-trades",
        type=Path,
        default=_LOGS / "overall_portfolio_trade_audits/engine_vxx_1pct_2016_2026/vxx_portfolio_bear_call.jsonl",
    )
    ap.add_argument(
        "--vxx-call-trades",
        type=Path,
        default=_LOGS / "overall_portfolio_trade_audits/engine_vxx_1pct_2016_2026/vxx_portfolio_long_call.jsonl",
    )
    args = ap.parse_args()

    cap = float(args.capital)
    sl = max(1, int(args.overlap_slice_contracts))
    prefix = str(args.out_prefix).strip() or "sharpe_opt_fullbook"
    out_dir = args.out_dir.expanduser().resolve()
    out_dir.mkdir(parents=True, exist_ok=True)

    cap_vrp = cap * float(args.capital_vrp_pct)
    cap_put = cap * float(args.capital_put_pct)
    cap_str = cap * float(args.capital_straddle_pct)
    cap_rr = cap * float(args.capital_risk_reversal_pct)
    cap_vxx = cap * float(args.capital_vxx_pct)
    bear_usd = cap_vxx * (float(args.vxx_bear_pct) / 100.0)
    call_usd = cap_vxx * (float(args.vxx_call_pct) / 100.0)

    put_p = args.put_trades.expanduser().resolve()
    str_p = args.straddle_trades.expanduser().resolve()
    rr_p = args.risk_reversal_trades.expanduser().resolve()
    vb_p = args.vxx_bear_trades.expanduser().resolve()
    vc_p = args.vxx_call_trades.expanduser().resolve()

    ref_put = _overlay_sleeve_risk_reference_usd(put_p, _infer_broker_risk_overlay)
    ref_str = _overlay_sleeve_risk_reference_usd(str_p, _infer_broker_risk_overlay)
    ref_rr = _overlay_sleeve_risk_reference_usd(rr_p, _infer_broker_risk_overlay)
    ref_vb = _overlay_sleeve_risk_reference_usd(vb_p, _infer_broker_risk_vxx)
    ref_vc = _overlay_sleeve_risk_reference_usd(vc_p, _infer_broker_risk_vxx)

    # --- VRP engine ---
    theta_dir = args.theta_dir.expanduser().resolve()
    d0, d1 = theta_chunks_date_bounds(theta_dir)
    yf_start = (d0 - pd.Timedelta(days=400)).strftime("%Y-%m-%d")
    yf_end = (d1 + pd.Timedelta(days=14)).strftime("%Y-%m-%d")
    spy_wide = normalize_spy_df(load_spy_vix_from_yfinance(yf_start, yf_end))
    ld = ThetaChunksLoader(theta_dir, spy_df=spy_wide)
    days = trading_days_intersecting_spy(ld, spy_wide.index, d0, d1)
    if str(args.start).strip():
        t0 = pd.Timestamp(str(args.start).strip())
        days = [d for d in days if d >= t0]
    if str(args.end).strip():
        t1 = pd.Timestamp(str(args.end).strip())
        days = [d for d in days if d <= t1]
    if int(args.max_days) > 0:
        days = days[: int(args.max_days)]
    if not days:
        raise SystemExit("No trading days after filters.")

    cfg = load_strategy_config_file(DEFAULT_STRATEGY_CONFIG_PATH)
    apply_strategy_params_to_vrp_backtester_module(cfg.strategy_params)

    bt = VRPBacktester(
        ld,
        initial_capital=cap,
        spy_df=spy_wide,
        vol_risk_scaling=False,
        r2_crossover_filters=True,
        dd_risk_scaling=True,
        sleeve_risk_fractions=dict(cfg.sleeve_risk_fractions),
        overlap_portfolio=True,
        overlap_slice_contracts=sl,
    )
    if bool(args.no_progress):
        print(f"Starting VRP backtest ({len(days)} sessions) …", flush=True)
    bt.run_backtest(trading_days=days, show_progress=not bool(args.no_progress))

    vrp_jsonl = (out_dir / f"{prefix}_vrp_trades.jsonl").resolve()
    vrp_csv = (out_dir / f"{prefix}_vrp_trade_log_with_legs.csv").resolve()
    _write_vrp_jsonl(vrp_jsonl, bt.trade_log)
    _write_vrp_csv(vrp_csv, bt.trade_log)

    vrp_scale = cap_vrp / max(float(args.vrp_ref), 1.0)
    vrp_flat: list[dict] = []
    for t in bt.trade_log:
        raw = float(t.pnl_usd)
        vrp_flat.append(
            {
                "sleeve": "vrp",
                "entry_date": pd.Timestamp(t.entry_date).strftime("%Y-%m-%d"),
                "exit_date": pd.Timestamp(t.exit_date).strftime("%Y-%m-%d"),
                "pnl_usd_scaled": round(raw * vrp_scale, 6),
                "pnl_usd_raw_in_log": round(raw, 6),
                "sleeve_scale_factor": round(vrp_scale, 10),
                "sleeve_risk_ref_usd": round(float(args.vrp_ref), 2),
                "sleeve_capital_usd": round(cap_vrp, 2),
                "regime": str(t.regime),
                "qty": int(t.qty),
                "exit_reason": str(t.exit_reason),
                "legs_json": str(getattr(t, "legs_json", "") or ""),
            }
        )

    rows_put = _load_jsonl(put_p) if put_p.is_file() else []
    rows_str = _load_jsonl(str_p) if str_p.is_file() else []
    rows_rr = _load_jsonl(rr_p) if rr_p.is_file() else []
    rows_vb = _load_jsonl(vb_p) if vb_p.is_file() else []
    rows_vc = _load_jsonl(vc_p) if vc_p.is_file() else []

    scaled_put = _scale_rows(rows_put, sleeve="iv_otm_put", pnl_key="pnl_total", capital_usd=cap_put, ref_usd=ref_put)
    scaled_str = _scale_rows(rows_str, sleeve="iv_straddle", pnl_key="pnl_total", capital_usd=cap_str, ref_usd=ref_str)
    scaled_rr = _scale_rows(rows_rr, sleeve="iv_risk_reversal", pnl_key="pnl_total", capital_usd=cap_rr, ref_usd=ref_rr)
    scaled_vb = _scale_rows(rows_vb, sleeve="vxx_bear_call", pnl_key="pnl_total", capital_usd=bear_usd, ref_usd=ref_vb)
    scaled_vc = _scale_rows(rows_vc, sleeve="vxx_long_call", pnl_key="pnl_total", capital_usd=call_usd, ref_usd=ref_vc)

    def _write_simple_csv(path: Path, fieldnames: list[str], rows: list[dict]) -> None:
        path.parent.mkdir(parents=True, exist_ok=True)
        with path.open("w", newline="", encoding="utf-8") as f:
            w = csv.DictWriter(f, fieldnames=fieldnames, extrasaction="ignore")
            w.writeheader()
            for r in rows:
                w.writerow(r)

    base_cols = [
        "sleeve",
        "entry_date",
        "exit_date",
        "pnl_usd_scaled",
        "pnl_usd_raw_in_log",
        "sleeve_scale_factor",
        "sleeve_risk_ref_usd",
        "sleeve_capital_usd",
    ]
    _write_simple_csv(out_dir / f"{prefix}_iv_put_scaled.csv", base_cols, scaled_put)
    _write_simple_csv(out_dir / f"{prefix}_iv_straddle_scaled.csv", base_cols, scaled_str)
    _write_simple_csv(out_dir / f"{prefix}_iv_rr_scaled.csv", base_cols, scaled_rr)
    _write_simple_csv(out_dir / f"{prefix}_vxx_bear_scaled.csv", base_cols, scaled_vb)
    _write_simple_csv(out_dir / f"{prefix}_vxx_long_scaled.csv", base_cols, scaled_vc)

    combined = vrp_flat + scaled_put + scaled_str + scaled_rr + scaled_vb + scaled_vc

    def _sort_key(r: dict) -> tuple:
        return (r.get("exit_date") or "", r.get("sleeve") or "")

    combined.sort(key=_sort_key)
    comb_fields = list(vrp_flat[0].keys()) if vrp_flat else list(combined[0].keys())
    for r in combined:
        for k in comb_fields:
            if k not in r:
                r[k] = ""
    _write_simple_csv(out_dir / f"{prefix}_ALL_SLEEVES_trade_log.csv", comb_fields, combined)

    eq_csv = (out_dir / f"{prefix}_equity_daily.csv").resolve()
    frame = execute_portfolio_merge(
        vrp_trades=vrp_jsonl,
        total_portfolio_capital=cap,
        capital_vrp=cap_vrp,
        capital_put=cap_put,
        capital_straddle=cap_str,
        capital_risk_reversal=cap_rr,
        capital_vxx=cap_vxx,
        vxx_bear_pct=float(args.vxx_bear_pct),
        vxx_call_pct=float(args.vxx_call_pct),
        vrp_ref=float(args.vrp_ref),
        put_trades=put_p,
        straddle_trades=str_p,
        risk_reversal_trades=rr_p,
        vxx_bear_trades=vb_p,
        vxx_call_trades=vc_p,
        out_csv=eq_csv,
        print_report=True,
        print_vxx_sweep=False,
    )
    m_full = _metrics_block(frame["eq_full_portfolio"], "full")

    manifest = out_dir / f"{prefix}_manifest.txt"
    lines = [
        "Sharpe-opt full-book trade export (VRP re-run + scaled overlay JSONLs)",
        f"prefix: {prefix}",
        f"window: {pd.Timestamp(days[0]).date()} → {pd.Timestamp(days[-1]).date()}  n_sessions={len(days)}",
        f"VRP preset: overlap_slice_contracts={sl} vol_off r2_on dd_on overlap_portfolio=True",
        f"capital={cap:g}  capital_vrp={cap_vrp:g}  vrp_ref={float(args.vrp_ref):g}",
        "",
        "Merge sleeve budgets (USD) and risk refs (JSONL sum broker_risk when documented):",
        f"  put:     capital={cap_put:,.2f}  ref={ref_put:,.2f}  scale={cap_put/ref_put:.8f}",
        f"  straddle: capital={cap_str:,.2f}  ref={ref_str:,.2f}  scale={cap_str/ref_str:.8f}",
        f"  rr:      capital={cap_rr:,.2f}  ref={ref_rr:,.2f}  scale={cap_rr/ref_rr:.8f}",
        f"  vxx total: capital={cap_vxx:,.2f}  bear_usd={bear_usd:,.2f} ref={ref_vb:,.2f} scale={bear_usd/ref_vb:.8f}",
        f"             call_usd={call_usd:,.2f} ref={ref_vc:,.2f} scale={call_usd/ref_vc:.8f}",
        "",
        "Trade counts:",
        f"  vrp_closed={len(bt.trade_log)}  iv_put={len(scaled_put)}  iv_straddle={len(scaled_str)}  iv_rr={len(scaled_rr)}  vxx_bear={len(scaled_vb)}  vxx_long={len(scaled_vc)}  combined_rows={len(combined)}",
        "",
        "Merged full-book metrics (execute_portfolio_merge):",
        f"  return%={m_full['return_pct']:.2f}  maxDD%={m_full['max_dd_pct']:.2f}  CAGR%={m_full['cagr_pct']:.2f}  Sharpe={m_full['sharpe']:.3f}",
        "",
        "Outputs:",
        f"  {vrp_jsonl}",
        f"  {vrp_csv}",
        f"  {out_dir / (prefix + '_iv_put_scaled.csv')}",
        f"  {out_dir / (prefix + '_iv_straddle_scaled.csv')}",
        f"  {out_dir / (prefix + '_iv_rr_scaled.csv')}",
        f"  {out_dir / (prefix + '_vxx_bear_scaled.csv')}",
        f"  {out_dir / (prefix + '_vxx_long_scaled.csv')}",
        f"  {out_dir / (prefix + '_ALL_SLEEVES_trade_log.csv')}",
        f"  {eq_csv}",
        "",
        "Source JSONLs (unscaled; PnL scaled in export):",
        f"  put: {put_p}",
        f"  straddle: {str_p}",
        f"  rr: {rr_p}",
        f"  vxx_bear: {vb_p}",
        f"  vxx_call: {vc_p}",
        "",
        "Reproduce:",
        f"  PYTHONUNBUFFERED=1 .venv/bin/python RenTech/strategy_stack/export_sharpe_opt_fullbook_trades.py \\",
        f"    --overlap-slice-contracts {sl} --capital {cap:g} --capital-vrp-pct {args.capital_vrp_pct:g} \\",
        f"    --capital-put-pct {args.capital_put_pct:.8f} --capital-straddle-pct {args.capital_straddle_pct:.8f} \\",
        f"    --capital-risk-reversal-pct {args.capital_risk_reversal_pct:.8f} --capital-vxx-pct {args.capital_vxx_pct:.8f} \\",
        f"    --out-prefix {prefix}",
    ]
    manifest.write_text("\n".join(lines) + "\n", encoding="utf-8")

    print(f"\nWrote manifest: {manifest}")
    print(f"Combined trade log: {out_dir / (prefix + '_ALL_SLEEVES_trade_log.csv')}")


if __name__ == "__main__":
    main()
