#!/usr/bin/env python3
"""
Batch runner for the **diverse_theta_strategies_v1** catalog (one module per strategy).

For each strategy module ``strategies.s000`` … ``strategies.s099`` the runner:

1. Builds Theta context via ``prepare_theta_research_context`` (chains + IV + skew caches).
2. Augments the SPY/VIX panel with extra columns (see :mod:`panel`).
3. Wraps each strategy's ``wants_entry`` in the sequential backtest engine used elsewhere
   in the repo (no overlapping positions; hold in **sessions**).

Output JSON contains per-strategy Sharpe, trade count, and metadata copied from each module.

Run::

    cd /Users/robzingale/trading_bot && .venv/bin/python -m RenTech.strategy_stack.diverse_theta_strategies_v1.runner \\
        --start 2021-01-04 --end 2024-12-31 \\
        --out-json RenTech/data/logs/diverse_theta_strategies_v1_2021_2024.json
"""
from __future__ import annotations

import argparse
import importlib
import json
import math
import sys
import time
from pathlib import Path
from typing import Any, Callable

import numpy as np
import pandas as pd

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

from RenTech.core.options_data_loader import OptionChain
from RenTech.strategy_stack import research_literature_theta_strategies as L
from RenTech.strategy_stack.diverse_theta_strategies_v1.context import ResearchContext
from RenTech.strategy_stack.diverse_theta_strategies_v1.panel import augment_research_panel

SignalFn = Callable[[int, pd.Series, OptionChain, float], bool]


def _trade_fn(kind: str, tp: tuple) -> L.TradeFn:
    k = str(kind).strip().lower()
    if k == "ss":
        return L._wrap_straddle_short(tp)
    if k == "sl":
        return L._wrap_straddle_long(tp)
    if k == "sg":
        return L._wrap_strangle(tp)
    if k == "put":
        return L._wrap_put(tp)
    if k == "rr":
        return L._wrap_rr(tp)
    if k == "vert":
        return L._wrap_vert(tp)
    if k == "vtc":
        return L._wrap_vert_call(tp)
    raise ValueError(f"Unknown TRADE_KIND {kind!r}")


def _load_strategy_module(index: int):
    name = f"RenTech.strategy_stack.diverse_theta_strategies_v1.strategies.s{index:03d}"
    return importlib.import_module(name)


def main() -> None:
    ap = argparse.ArgumentParser(description="Run diverse_theta_strategies_v1 catalog (D000–D099)")
    ap.add_argument("--theta-dir", type=Path, default=L._DEFAULT_THETA)
    ap.add_argument("--capital", type=float, default=1_000_000.0)
    ap.add_argument("--start", type=str, default="")
    ap.add_argument("--end", type=str, default="")
    ap.add_argument("--max-days", type=int, default=0)
    ap.add_argument(
        "--out-json",
        type=Path,
        default=_REPO_ROOT / "RenTech" / "data" / "logs" / "diverse_theta_strategies_v1_batch.json",
    )
    ap.add_argument("--from-index", type=int, default=0, help="Inclusive strategy index 0–99")
    ap.add_argument("--to-index", type=int, default=99, help="Inclusive strategy index 0–99")
    args = ap.parse_args()

    t0 = time.perf_counter()
    days, panel0, get_chain, iv_atm, skew, n_contracts, spy_wide = L.prepare_theta_research_context(
        theta_dir=args.theta_dir.expanduser().resolve(),
        capital=float(args.capital),
        start=str(args.start).strip(),
        end=str(args.end).strip(),
        max_days=int(args.max_days),
    )
    panel = augment_research_panel(panel0)
    ctx = ResearchContext(days, panel, get_chain, iv_atm, skew, n_contracts)
    print(
        f"Context ready: {len(days)} sessions in {(time.perf_counter() - t0)/60:.2f} min",
        flush=True,
    )

    lo = max(0, int(args.from_index))
    hi = min(99, int(args.to_index))
    results: list[dict[str, Any]] = []
    daily_ret: dict[str, pd.Series] = {}

    for k in range(lo, hi + 1):
        mod = _load_strategy_module(k)
        meta = dict(getattr(mod, "META"))
        sid = str(meta.get("sid", f"D{k:03d}"))
        hold = int(getattr(mod, "HOLD_SESSIONS"))
        tk = str(getattr(mod, "TRADE_KIND"))
        tp = tuple(getattr(mod, "TRADE_PARAMS"))
        tfn = _trade_fn(tk, tp)
        wants = getattr(mod, "wants_entry")

        def _sig(i: int, row: pd.Series, ch: OptionChain, spy: float) -> bool:
            return bool(wants(i, row, ch, spy, ctx))

        ex, pnls, ntr = L.run_signal_backtest(days, get_chain, panel, _sig, hold, tfn, tp)
        eq, sh = L.equity_curve_from_realized(ex, pnls, days, float(args.capital))
        dret = eq.pct_change().replace([np.inf, -np.inf], np.nan).fillna(0.0)
        daily_ret[sid] = dret
        results.append(
            {
                "sid": sid,
                "theme": meta.get("theme", ""),
                "title": meta.get("title", ""),
                "module": f"s{k:03d}",
                "hold_sessions": hold,
                "trade_kind": tk,
                "trade_params": list(tp),
                "trades": int(ntr),
                "sharpe": float(sh) if math.isfinite(sh) else float("nan"),
            }
        )
        print(
            f"{sid} {str(meta.get('title', ''))[:48]!r} trades={ntr} sharpe={sh:.3f}",
            flush=True,
        )

    df_ret = pd.DataFrame(daily_ret).reindex(pd.DatetimeIndex([L._norm(d) for d in days]))
    n_sess = len(days)
    min_p = max(5, min(50, max(3, n_sess - 2)))
    corr = df_ret.corr(method="pearson", min_periods=min_p)

    out_path = args.out_json.expanduser().resolve()
    out_path.parent.mkdir(parents=True, exist_ok=True)
    payload = {
        "meta": {
            "catalog": "diverse_theta_strategies_v1",
            "sessions": len(days),
            "first_day": str(days[0].date()) if days else "",
            "last_day": str(days[-1].date()) if days else "",
            "capital": float(args.capital),
            "from_index": lo,
            "to_index": hi,
        },
        "results": results,
        "correlation_matrix": json.loads(corr.to_json()) if len(corr) else {},
    }
    out_path.write_text(json.dumps(payload, indent=2), encoding="utf-8")
    print(f"Wrote {out_path}", flush=True)


if __name__ == "__main__":
    main()
