#!/usr/bin/env python3
"""
Batch-evaluate the **100-strategy catalog**, compute Sharpe + daily PnL correlations,
greedy-select a **low-correlation** high-Sharpe subset, and print a **critique** loop
that suggests expanding underrepresented / weak sleeves until targets are met.

Run::

    cd /Users/robzingale/trading_bot && .venv/bin/python RenTech/strategy_stack/literature_search_agent.py \\
        --start 2016-01-04 --end 2022-12-31 --min-trades 25 --target-sharpe 1.5 \\
        --max-pairwise-corr 0.45 --min-uncorrelated 12

Slice (e.g. stress window only; ``command_meta`` records ``start_arg`` / ``end_arg``)::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 .venv/bin/python RenTech/strategy_stack/literature_search_agent.py \\
        --start 2022-01-04 --end 2024-12-31 --min-trades 18 \\
        --out-json RenTech/data/logs/literature_search_agent_batch_2022_2024.json

Default greedy portfolio caps (override with ``--family-cap FAMILY=N``, disable with ``--no-family-caps``):
``short_rr=2``, ``putwrite=2``. JSON output includes ``selection.greedy_uncapped`` vs ``selection.greedy_capped``.
"""
from __future__ import annotations

import argparse
import json
import math
import sys
import time
from pathlib import Path
from typing import Any

import numpy as np
import pandas as pd

_REPO_ROOT = Path(__file__).resolve().parents[2]
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.literature_strategy_catalog import StrategySpec, build_catalog_100
from RenTech.strategy_stack import research_literature_theta_strategies as L

_DEFAULT_OUT = _REPO_ROOT / "RenTech" / "data" / "logs" / "literature_search_agent_batch.json"
_DEFAULT_THETA = _REPO_ROOT / "RenTech" / "data" / "theta_chunks"

# Greedy low-correlation portfolio: limit how many passers from dominant empirical families.
_DEFAULT_FAMILY_CAPS: dict[str, int] = {"short_rr": 2, "putwrite": 2}


def _merge_family_caps(no_family_caps: bool, overrides: list[str]) -> dict[str, int]:
    if no_family_caps:
        base: dict[str, int] = {}
    else:
        base = dict(_DEFAULT_FAMILY_CAPS)
    for item in overrides:
        if "=" not in item:
            raise SystemExit(f"--family-cap expects FAMILY=N, got {item!r}")
        k, v = item.split("=", 1)
        base[k.strip()] = int(v.strip())
    return base


def _compile_signal(
    spec: StrategySpec,
    panel: pd.DataFrame,
    iv_atm: dict[tuple[int, int], float | None],
    skew_put_minus_call_iv: dict[tuple[int, int, float, float], float | None],
    n_contracts: list[int],
    spy_wide: pd.DataFrame,
) -> L.SignalFn:
    sk = spec.sig_kind
    sp = spec.sig_params

    if sk == "vrp":
        thr, target_dte = float(sp[0]), int(sp[1])

        def sig(i: int, row: pd.Series, ch: OptionChain, _s: float) -> bool:
            iv = iv_atm.get((i, target_dte))
            rv = float(row["rv21"]) if pd.notna(row.get("rv21")) else float("nan")
            if iv is None or not math.isfinite(rv) or rv <= 0:
                return False
            return iv > rv + thr * rv

        return sig

    if sk == "lowvix":
        vlim, pdel, cdel = float(sp[0]), float(sp[1]), float(sp[2])

        def sig(i: int, row: pd.Series, ch: OptionChain, _s: float) -> bool:
            return float(row["vix_close"]) < vlim and n_contracts[i] > 80

        return sig

    if sk == "spike":
        (pct,) = (float(sp[0]),)

        def sig(i: int, row: pd.Series, ch: OptionChain, _s: float) -> bool:
            if i < 5:
                return False
            v0 = float(panel.iloc[i - 5]["vix_close"])
            v1 = float(row["vix_close"])
            return v1 > v0 * (1.0 + pct)

        return sig

    if sk == "vixfall":
        (dvx,) = (float(sp[0]),)

        def sig(i: int, row: pd.Series, ch: OptionChain, _s: float) -> bool:
            if i < 5:
                return False
            return float(row["vix_close"]) - float(panel.iloc[i - 5]["vix_close"]) < dvx

        return sig

    if sk == "putwrite":
        delt, use50, vix_cap = float(sp[0]), bool(sp[1]), float(sp[2])

        def sig(i: int, row: pd.Series, ch: OptionChain, spy: float) -> bool:
            sma_col = "sma_50" if use50 else "sma_200"
            sma = float(row[sma_col]) if pd.notna(row.get(sma_col)) else float("nan")
            return math.isfinite(sma) and spy > sma and float(row["vix_close"]) < vix_cap

        return sig

    if sk == "rr":
        gap_iv, target_dte, put_d, call_d = float(sp[0]), int(sp[1]), float(sp[2]), float(sp[3])

        def sig(i: int, row: pd.Series, ch: OptionChain, _s: float) -> bool:
            diff = skew_put_minus_call_iv.get((i, target_dte, put_d, call_d))
            return diff is not None and diff > gap_iv

        return sig

    if sk == "vvix":
        vvix_vix_cap, target_dte = float(sp[0]), int(sp[1])

        def sig(i: int, row: pd.Series, ch: OptionChain, _s: float) -> bool:
            if "vvix_close" not in spy_wide.columns:
                return False
            vv = row.get("vvix_close")
            vx = row.get("vix_close")
            if pd.isna(vv) or pd.isna(vx) or float(vx) <= 0:
                return False
            if float(vv) / float(vx) > vvix_vix_cap:
                return False
            iv = iv_atm.get((i, target_dte))
            rv = float(row["rv21"]) if pd.notna(row.get("rv21")) else float("nan")
            return iv is not None and math.isfinite(rv) and rv > 0 and iv > rv * 1.05

        return sig

    if sk == "rvcomp":
        rv_gap, target_dte = float(sp[0]), int(sp[1])

        def sig(i: int, row: pd.Series, ch: OptionChain, _s: float) -> bool:
            r5 = float(row["rv5"]) if pd.notna(row.get("rv5")) else float("nan")
            r21 = float(row["rv21"]) if pd.notna(row.get("rv21")) else float("nan")
            if not (math.isfinite(r5) and math.isfinite(r21)) or r21 <= 0:
                return False
            if r21 - r5 <= rv_gap:
                return False
            iv = iv_atm.get((i, target_dte))
            return iv is not None and iv > r21

        return sig

    if sk == "midvx_vert":

        def sig(i: int, row: pd.Series, ch: OptionChain, _s: float) -> bool:
            v = float(row["vix_close"])
            return 14.0 < v < 22.0

        return sig

    if sk == "hivix":
        vix_floor, rv_mult, target_dte = float(sp[0]), float(sp[1]), int(sp[2])

        def sig(i: int, row: pd.Series, ch: OptionChain, _s: float) -> bool:
            if float(row["vix_close"]) < vix_floor:
                return False
            iv = iv_atm.get((i, target_dte))
            rv = float(row["rv21"]) if pd.notna(row.get("rv21")) else float("nan")
            if iv is None or not math.isfinite(rv) or rv <= 0:
                return False
            return iv > rv + rv_mult * rv

        return sig

    if sk == "lowiv":
        vix_cap, iv_rv_mult, target_dte = float(sp[0]), float(sp[1]), int(sp[2])

        def sig(i: int, row: pd.Series, ch: OptionChain, _s: float) -> bool:
            if float(row["vix_close"]) > vix_cap:
                return False
            iv = iv_atm.get((i, target_dte))
            rv = float(row["rv21"]) if pd.notna(row.get("rv21")) else float("nan")
            if iv is None or not math.isfinite(rv) or rv <= 0:
                return False
            return iv < rv * iv_rv_mult

        return sig

    if sk == "downtrend":
        use50, vix_min, vix_max = bool(sp[0]), float(sp[1]), float(sp[2])

        def sig(i: int, row: pd.Series, ch: OptionChain, spy: float) -> bool:
            sma_col = "sma_50" if use50 else "sma_200"
            sma = float(row[sma_col]) if pd.notna(row.get(sma_col)) else float("nan")
            if not math.isfinite(sma):
                return False
            vix = float(row["vix_close"])
            return spy < sma and (vix_min <= vix <= vix_max)

        return sig

    raise ValueError(f"Unknown sig_kind {sk!r} for {spec.sid}")


def _compile_trade(spec: StrategySpec) -> L.TradeFn:
    tk = spec.trade_kind
    tp = spec.trade_params
    if tk == "ss":
        return L._wrap_straddle_short(tp)
    if tk == "sl":
        return L._wrap_straddle_long(tp)
    if tk == "sg":
        return L._wrap_strangle(tp)
    if tk == "rr":
        return L._wrap_rr(tp)
    if tk == "put":
        return L._wrap_put(tp)
    if tk == "vert":
        return L._wrap_vert(tp)
    if tk == "vtc":
        return L._wrap_vert_call(tp)
    raise ValueError(f"Unknown trade_kind {tk!r}")


def greedy_low_corr_subset(
    ids: list[str],
    corr: pd.DataFrame,
    sharpe_map: dict[str, float],
    max_corr: float,
    *,
    family_by_sid: dict[str, str],
    family_cap: dict[str, int] | None = None,
) -> list[str]:
    """Pick strategies greedily: highest Sharpe first; keep if max |corr| to chosen < threshold.

    If ``family_cap`` is set, each listed family contributes at most that many strategies.
    Families omitted from the dict are unlimited.
    """

    def _under_cap(fam: str, counts: dict[str, int]) -> bool:
        if not family_cap:
            return True
        lim = family_cap.get(fam)
        if lim is None:
            return True
        return counts.get(fam, 0) < lim

    ordered = sorted(ids, key=lambda k: sharpe_map.get(k, float("-inf")), reverse=True)
    chosen: list[str] = []
    fam_counts: dict[str, int] = {}
    for sid in ordered:
        fam = family_by_sid.get(sid, "")
        if not _under_cap(fam, fam_counts):
            continue
        if not chosen:
            chosen.append(sid)
            fam_counts[fam] = fam_counts.get(fam, 0) + 1
            continue
        ok = True
        for cj in chosen:
            rho = abs(float(corr.loc[sid, cj]))
            if rho > max_corr:
                ok = False
                break
        if ok:
            chosen.append(sid)
            fam_counts[fam] = fam_counts.get(fam, 0) + 1
    return chosen


def critique_report(
    results: list[dict[str, Any]],
    corr: pd.DataFrame | None,
    *,
    target_sharpe: float,
    min_trades: int,
    max_pairwise_corr: float,
    min_uncorrelated: int,
    family_caps: dict[str, int],
    no_family_caps: bool,
) -> tuple[list[str], dict[str, Any]]:
    lines: list[str] = []
    selection: dict[str, Any] = {}
    sharp = [r["sharpe"] for r in results if math.isfinite(r["sharpe"])]
    passers = [r for r in results if r["trades"] >= min_trades and r["sharpe"] >= target_sharpe]
    lines.append(
        f"Pass rate: {len(passers)}/{len(results)} with Sharpe≥{target_sharpe} and trades≥{min_trades}."
    )
    if sharp:
        lines.append(f"Sharpe distribution: min={min(sharp):.3f} median={float(np.median(sharp)):.3f} max={max(sharp):.3f}")

    by_family: dict[str, list[float]] = {}
    for r in results:
        by_family.setdefault(r["family"], []).append(float(r["sharpe"]))
    weak_families = sorted(
        by_family.keys(),
        key=lambda f: float(np.median(by_family[f])) if by_family[f] else -999,
    )[:5]
    lines.append(f"Lowest median-Sharpe families (consider richer grids / new signals): {weak_families}")

    if corr is not None and len(corr) >= 2:
        triu = np.triu(np.ones(corr.shape), k=1).astype(bool)
        vals = corr.where(triu).stack().dropna()
        vals = vals[np.isfinite(vals)]
        if len(vals):
            lines.append(
                f"Pairwise |corr| of daily equity returns: median={float(np.median(np.abs(vals))):.3f} "
                f"p90={float(np.quantile(np.abs(vals), 0.9)):.3f} max={float(np.max(np.abs(vals))):.3f}"
            )
            high_pairs = sum(bool(abs(float(v)) > max_pairwise_corr) for v in vals)
            lines.append(
                f"Pairs with |ρ|>{max_pairwise_corr}: {high_pairs} / {len(vals)} — "
                "high counts imply redundant sleeves (same vol regime)."
            )
        else:
            lines.append(
                "Correlation matrix has no finite pairwise entries (short window or sparse overlap); "
                "rerun with more sessions for overlap stats."
            )

    sharpe_map = {r["sid"]: float(r["sharpe"]) for r in results}
    sid_family = {r["sid"]: r["family"] for r in results}
    good_ids = [r["sid"] for r in passers]
    if corr is not None and good_ids:
        sub = corr.loc[good_ids, good_ids]
        greedy_u = greedy_low_corr_subset(
            good_ids,
            sub,
            sharpe_map,
            max_pairwise_corr,
            family_by_sid=sid_family,
            family_cap=None,
        )
        selection["greedy_uncapped"] = greedy_u
        lines.append(
            f"Greedy uncorrelated subset (no family cap, |ρ|≤{max_pairwise_corr}): size={len(greedy_u)} "
            f"(target min {min_uncorrelated})."
        )

        if no_family_caps:
            selection["greedy_capped"] = greedy_u
            selection["family_caps_applied"] = {}
            shortfall = len(greedy_u) < min_uncorrelated
        else:
            selection["family_caps_applied"] = dict(family_caps)
            greedy_c = greedy_low_corr_subset(
                good_ids,
                sub,
                sharpe_map,
                max_pairwise_corr,
                family_by_sid=sid_family,
                family_cap=family_caps,
            )
            selection["greedy_capped"] = greedy_c
            lines.append(
                f"Greedy uncorrelated subset (family caps {family_caps}, |ρ|≤{max_pairwise_corr}): "
                f"size={len(greedy_c)} (target min {min_uncorrelated})."
            )
            shortfall = len(greedy_c) < min_uncorrelated

        if shortfall:
            lines.append(
                "CRITIQUE: Increase diversification — add sleeves with different Greeks "
                "(e.g. bear-call / long-gamma) or signals keyed to **spot** not just VIX; "
                "widen date window; relax single-position constraint for orthogonal sleeves only."
            )

    return lines, selection


def main() -> None:
    ap = argparse.ArgumentParser(description="100-strategy literature batch + correlation critique")
    ap.add_argument("--theta-dir", type=Path, default=_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("--min-trades", type=int, default=25)
    ap.add_argument("--target-sharpe", type=float, default=1.5)
    ap.add_argument("--max-pairwise-corr", type=float, default=0.45)
    ap.add_argument("--min-uncorrelated", type=int, default=12)
    ap.add_argument(
        "--family-cap",
        action="append",
        default=[],
        metavar="FAMILY=N",
        help="Per-family max count in capped greedy selection (repeatable). Defaults: short_rr=2, putwrite=2.",
    )
    ap.add_argument(
        "--no-family-caps",
        action="store_true",
        help="Disable per-family caps (capped subset equals uncapped).",
    )
    ap.add_argument("--out-json", type=Path, default=_DEFAULT_OUT)
    args = ap.parse_args()

    family_caps = _merge_family_caps(bool(args.no_family_caps), list(args.family_cap))

    catalog = build_catalog_100()
    t0 = time.perf_counter()
    print(f"Preparing Theta context ({len(catalog)} strategies)…", flush=True)
    days, panel, get_chain, iv_atm, skew_put_minus_call_iv, n_contracts, spy_wide = L.prepare_theta_research_context(
        theta_dir=args.theta_dir,
        capital=float(args.capital),
        start=str(args.start),
        end=str(args.end),
        max_days=int(args.max_days),
    )
    print(f"Context ready in {(time.perf_counter()-t0)/60:.2f} min; {len(days)} sessions.", flush=True)

    results: list[dict[str, Any]] = []
    daily_ret_matrix: dict[str, pd.Series] = {}

    for spec in catalog:
        sig = _compile_signal(spec, panel, iv_atm, skew_put_minus_call_iv, n_contracts, spy_wide)
        tfn = _compile_trade(spec)
        ex, pnls, ntr = L.run_signal_backtest(
            days, get_chain, panel, sig, int(spec.hold), tfn, spec.trade_params
        )
        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_matrix[spec.sid] = dret
        results.append(
            {
                "sid": spec.sid,
                "family": spec.family,
                "description": spec.description,
                "hold": spec.hold,
                "trades": int(ntr),
                "sharpe": float(sh) if math.isfinite(sh) else float("nan"),
                "sig_kind": spec.sig_kind,
                "trade_kind": spec.trade_kind,
            }
        )

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

    out_path = args.out_json.expanduser().resolve()
    out_path.parent.mkdir(parents=True, exist_ok=True)
    critique_lines, selection_payload = critique_report(
        results,
        corr,
        target_sharpe=float(args.target_sharpe),
        min_trades=int(args.min_trades),
        max_pairwise_corr=float(args.max_pairwise_corr),
        min_uncorrelated=int(args.min_uncorrelated),
        family_caps=family_caps,
        no_family_caps=bool(args.no_family_caps),
    )
    payload = {
        "command_meta": {
            "script": "literature_search_agent.py",
            "n_strategies": len(catalog),
            "sessions": len(days),
            "start_arg": str(args.start).strip(),
            "end_arg": str(args.end).strip(),
            "max_days_arg": int(args.max_days),
            "first_day": str(days[0].date()) if days else "",
            "last_day": str(days[-1].date()) if days else "",
            "capital": float(args.capital),
            "min_trades": int(args.min_trades),
            "target_sharpe": float(args.target_sharpe),
            "family_caps_default": dict(_DEFAULT_FAMILY_CAPS),
            "family_caps_effective": {} if args.no_family_caps else dict(family_caps),
            "no_family_caps": bool(args.no_family_caps),
        },
        "results": results,
        "selection": selection_payload,
        "critique": critique_lines,
    }
    out_path.write_text(json.dumps(payload, indent=2), encoding="utf-8")
    print(json.dumps(payload["critique"], indent=2))
    print(f"\nWrote {out_path}", flush=True)


if __name__ == "__main__":
    main()
