#!/usr/bin/env python3
"""
Grid-search optimizer for Markov VIX futures L/S with **fixed-expiry** execution.

Precomputes Markov features once per VIX-regime setting, then sweeps trading
rules and risk limits.  Ranks configs on train-window Sharpe with drawdown
filters, then reports out-of-sample validation metrics.

Example::

    cd /Users/robzingale/trading_bot
    PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/optimize_markov_vix_futures_ls.py \\
        --train-start 2016-01-04 --train-end 2022-12-30 \\
        --valid-start 2023-01-01
"""

from __future__ import annotations

import argparse
import itertools
import json
import sys
from pathlib import Path

import numpy as np
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.markov_chain_trading import MarkovChainTradingModel
from RenTech.strategy_stack.run_markov_vix_futures_backtest import (
    CONTRACTS_PATH,
    _ensure_vix_futures_panel,
    _metrics,
    _prepare_futures_dict,
    assemble_vix_weights_from_features,
    build_markov_vix_feature_cache,
)
from RenTech.strategy_stack.vix_fixed_calendar_engine import VixContractStore
from RenTech.strategy_stack.vix_fixed_outright_engine import simulate_fixed_outright_portfolio

DEFAULT_OUT = _REPO / "RenTech" / "data" / "logs" / "markov_vix_futures_ls_opt"


def _param_key(params: dict) -> tuple:
    return (
        params["vix_regime"],
        params["allow_short"],
        params["edge_threshold"],
        params["kelly_mult"],
        params["max_short"],
        params["gross_cap"],
        params["rel_spread"],
        params["rel_threshold"],
    )


def _build_grid(full: bool, long_only_grid: bool) -> list[dict]:
    if full:
        edge = [0.02, 0.03, 0.04, 0.05, 0.06, 0.08]
        kelly = [0.15, 0.25, 0.35, 0.50]
        max_short = [0.25, 0.35, 0.50]
        gross = [0.75, 1.0, 1.25]
        rel_t = [0.04, 0.06, 0.08]
    else:
        edge = [0.03, 0.04, 0.05, 0.06, 0.08]
        kelly = [0.20, 0.30, 0.40]
        max_short = [0.30, 0.50]
        gross = [0.75, 1.0]
        rel_t = [0.05, 0.07]

    allow_opts = [False] if long_only_grid else [False, True]
    combos: list[dict] = []
    for vr, als, et, km, ms, gc, rs, rt in itertools.product(
        [False, True],
        allow_opts,
        edge,
        kelly,
        max_short,
        gross,
        [False, True],
        rel_t,
    ):
        if not rs and rt != rel_t[0]:
            continue
        if not als and ms != max_short[0]:
            continue
        combos.append(
            {
                "vix_regime": vr,
                "allow_short": als,
                "edge_threshold": et,
                "kelly_mult": km,
                "max_short": ms,
                "gross_cap": gc,
                "rel_spread": rs,
                "rel_threshold": rt,
            }
        )
    return combos


def _evaluate_config(
    *,
    store: VixContractStore,
    sim_dates: pd.DatetimeIndex,
    score_start: pd.Timestamp,
    score_end: pd.Timestamp | None,
    weights: pd.DataFrame,
    edges: pd.DataFrame,
    decisions: pd.DataFrame,
    contracts: list[str],
    capital: float,
    cash_yield: float,
    gross_cap: float,
) -> dict:
    port, _ = simulate_fixed_outright_portfolio(
        store=store,
        dates=sim_dates.sort_values(),
        target_weights=weights.reindex(sim_dates)[contracts],
        edge_df=edges.reindex(sim_dates)[contracts],
        decision_df=decisions.reindex(sim_dates)[contracts],
        contracts=contracts,
        capital=capital,
        cash_annual_yield=cash_yield,
        gross_cap=gross_cap,
        record_trades=False,
    )
    mask = port.index >= score_start
    if score_end is not None:
        mask &= port.index <= score_end
    r = port.loc[mask, "portfolio_bar_ret"].fillna(0.0)
    m = _metrics(r, capital)
    m["avg_gross"] = float(port.loc[mask, "gross_exposure"].mean())
    return m


def _score_train(m: dict, *, min_return_pct: float, max_dd_pct: float) -> float:
    if m["n_days"] < 100:
        return float("-inf")
    if m["total_return_pct"] < min_return_pct:
        return float("-inf")
    if m["max_drawdown_pct"] < max_dd_pct:
        return float("-inf")
    sharpe = m["sharpe"]
    if not np.isfinite(sharpe):
        return float("-inf")
    dd = abs(m["max_drawdown_pct"])
    cagr = m["cagr_pct"]
    calmar = cagr / dd if dd > 1e-6 else cagr
    return sharpe + 0.05 * calmar


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
    ap.add_argument("--train-start", default="2016-01-04")
    ap.add_argument("--train-end", default="2022-12-30")
    ap.add_argument("--valid-start", default="2023-01-01")
    ap.add_argument("--valid-end", default="")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--cash-yield", type=float, default=0.04)
    ap.add_argument("--n-sims", type=int, default=1500)
    ap.add_argument("--matrix-refresh", type=int, default=5)
    ap.add_argument("--min-train-return-pct", type=float, default=0.0)
    ap.add_argument("--max-train-dd-pct", type=float, default=-55.0)
    ap.add_argument("--top-n", type=int, default=15)
    ap.add_argument("--full-grid", action="store_true", help="Larger parameter grid (slow)")
    ap.add_argument("--long-only-grid", action="store_true", help="Search long-only configs only")
    ap.add_argument("--skip-train", action="store_true", help="Validate top rows from existing grid CSV")
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT)
    args = ap.parse_args()

    train_start = pd.Timestamp(args.train_start)
    train_end = pd.Timestamp(args.train_end)
    valid_start = pd.Timestamp(args.valid_start)
    valid_end = pd.Timestamp(args.valid_end) if args.valid_end.strip() else None
    full_end = valid_end if valid_end is not None else pd.Timestamp("2099-12-31")

    panel = _ensure_vix_futures_panel()
    futures_dict = _prepare_futures_dict(panel)
    vix = panel["vix_spot"].astype(np.float64)
    store = VixContractStore(pd.read_parquet(CONTRACTS_PATH))

    base_model = MarkovChainTradingModel(
        n_states=10,
        horizon=21,
        n_sims=args.n_sims,
        calibration_mode="equity_shrink",
    )

    print("Precomputing Markov features (pooled) …", flush=True)
    master, feat_pooled, contracts = build_markov_vix_feature_cache(
        futures_dict,
        model=base_model,
        percentile_window=252,
        lookback=252,
        matrix_refresh=args.matrix_refresh,
        vix=None,
    )
    print("Precomputing Markov features (VIX regime) …", flush=True)
    _, feat_regime, _ = build_markov_vix_feature_cache(
        futures_dict,
        model=base_model,
        percentile_window=252,
        lookback=252,
        matrix_refresh=args.matrix_refresh,
        vix=vix,
        vix_lo=15.0,
        vix_hi=25.0,
    )

    combos = _build_grid(args.full_grid, args.long_only_grid)
    prefix = args.out_prefix.expanduser().resolve()
    prefix.parent.mkdir(parents=True, exist_ok=True)
    grid_path = Path(f"{prefix}_grid.csv")

    weight_cache: dict[tuple, tuple[pd.DataFrame, pd.DataFrame, pd.DataFrame]] = {}
    sim_dates = master[master >= train_start - pd.Timedelta(days=45)]
    rows: list[dict] = []

    if args.skip_train and grid_path.is_file():
        print(f"Loading train grid from {grid_path} …", flush=True)
        df = pd.read_csv(grid_path)
    else:
        print(f"Evaluating {len(combos)} configs …", flush=True)
        for i, params in enumerate(combos):
            features = feat_regime if params["vix_regime"] else feat_pooled
            key = _param_key(params)
            if key not in weight_cache:
                weight_cache[key] = assemble_vix_weights_from_features(
                    master,
                    features,
                    contracts,
                    edge_threshold=params["edge_threshold"],
                    kelly_mult=params["kelly_mult"],
                    allow_short=params["allow_short"],
                    max_short=params["max_short"],
                    gross_cap=params["gross_cap"],
                rel_spread=params["rel_spread"],
                rel_threshold=params["rel_threshold"],
                long_only=not params["allow_short"],
            )

            train_m = _evaluate_config(
                store=store,
                sim_dates=sim_dates,
                score_start=train_start,
                score_end=train_end,
                weights=weight_cache[key][0],
                edges=weight_cache[key][1],
                decisions=weight_cache[key][2],
                contracts=contracts,
                capital=args.capital,
                cash_yield=args.cash_yield,
                gross_cap=params["gross_cap"],
            )
            score = _score_train(
                train_m,
                min_return_pct=args.min_train_return_pct,
                max_dd_pct=args.max_train_dd_pct,
            )
            rows.append({**params, **{f"train_{k}": v for k, v in train_m.items()}, "train_score": score})
            if (i + 1) % 50 == 0:
                print(f"  … {i + 1}/{len(combos)}", flush=True)

        df = pd.DataFrame(rows).sort_values("train_score", ascending=False)
        df.to_csv(grid_path, index=False)
        print(f"Wrote {grid_path}")

    passing = df[df["train_score"] > float("-inf")]
    print(f"\nConfigs passing filters: {len(passing)} / {len(df)}")

    top = passing.head(args.top_n).copy()
    valid_sim_dates = master[master >= train_start - pd.Timedelta(days=45)]
    valid_rows: list[dict] = []
    for _, row in top.iterrows():
        params = {
            k: row[k]
            for k in (
                "vix_regime",
                "allow_short",
                "edge_threshold",
                "kelly_mult",
                "max_short",
                "gross_cap",
                "rel_spread",
                "rel_threshold",
            )
        }
        weights, edges, decisions = weight_cache.get(_param_key(params)) or assemble_vix_weights_from_features(
            master,
            feat_regime if params["vix_regime"] else feat_pooled,
            contracts,
            edge_threshold=float(params["edge_threshold"]),
            kelly_mult=float(params["kelly_mult"]),
            allow_short=bool(params["allow_short"]),
            max_short=float(params["max_short"]),
            gross_cap=float(params["gross_cap"]),
            rel_spread=bool(params["rel_spread"]),
            rel_threshold=float(params["rel_threshold"]),
            long_only=not bool(params["allow_short"]),
        )
        if _param_key(params) not in weight_cache:
            weight_cache[_param_key(params)] = (weights, edges, decisions)
        valid_m = _evaluate_config(
            store=store,
            sim_dates=valid_sim_dates,
            score_start=valid_start,
            score_end=valid_end,
            weights=weights,
            edges=edges,
            decisions=decisions,
            contracts=contracts,
            capital=args.capital,
            cash_yield=args.cash_yield,
            gross_cap=params["gross_cap"],
        )
        full_m = _evaluate_config(
            store=store,
            sim_dates=valid_sim_dates,
            score_start=train_start,
            score_end=full_end if valid_end is None else valid_end,
            weights=weights,
            edges=edges,
            decisions=decisions,
            contracts=contracts,
            capital=args.capital,
            cash_yield=args.cash_yield,
            gross_cap=params["gross_cap"],
        )
        valid_rows.append(
            {
                **params,
                "train_sharpe": row["train_sharpe"],
                "train_cagr_pct": row["train_cagr_pct"],
                "train_total_return_pct": row["train_total_return_pct"],
                "train_max_drawdown_pct": row["train_max_drawdown_pct"],
                "train_score": row["train_score"],
                **{f"valid_{k}": v for k, v in valid_m.items()},
                **{f"full_{k}": v for k, v in full_m.items()},
            }
        )

    valid_df = pd.DataFrame(valid_rows)
    if not valid_df.empty:
        valid_df = valid_df.sort_values("full_sharpe", ascending=False)

    top_path = Path(f"{prefix}_top_validated.csv")
    best_path = Path(f"{prefix}_best.json")

    valid_df.to_csv(top_path, index=False)

    if not valid_df.empty:
        b = valid_df.iloc[0]
        params = {
            k: b[k]
            for k in (
                "vix_regime",
                "allow_short",
                "edge_threshold",
                "kelly_mult",
                "max_short",
                "gross_cap",
                "rel_spread",
                "rel_threshold",
            )
        }
        best_out = {
            "params": params,
            "train": {
                "sharpe": float(b["train_sharpe"]),
                "cagr_pct": float(b["train_cagr_pct"]),
                "total_return_pct": float(b["train_total_return_pct"]),
                "max_drawdown_pct": float(b["train_max_drawdown_pct"]),
            },
            "validation": {
                "sharpe": float(b["valid_sharpe"]),
                "cagr_pct": float(b["valid_cagr_pct"]),
                "total_return_pct": float(b["valid_total_return_pct"]),
                "max_drawdown_pct": float(b["valid_max_drawdown_pct"]),
            },
            "full_sample": {
                "sharpe": float(b["full_sharpe"]),
                "cagr_pct": float(b["full_cagr_pct"]),
                "total_return_pct": float(b["full_total_return_pct"]),
                "max_drawdown_pct": float(b["full_max_drawdown_pct"]),
            },
        }
        best_path.write_text(json.dumps(best_out, indent=2, default=str) + "\n")

    print(f"\nWrote {grid_path}")
    print(f"Wrote {top_path}")
    if not valid_df.empty:
        print(f"Wrote {best_path}")

    if valid_df.empty:
        print("\nNo configs passed train filters. Relax --min-train-return-pct or --max-train-dd-pct.")
        return

    print("\n=== Top configs (full-sample Sharpe, fixed-expiry) ===\n")
    show_cols = [
        "vix_regime", "allow_short", "edge_threshold", "kelly_mult", "max_short",
        "gross_cap", "rel_spread", "rel_threshold",
        "train_sharpe", "train_cagr_pct", "train_max_drawdown_pct",
        "valid_sharpe", "valid_cagr_pct", "valid_total_return_pct", "valid_max_drawdown_pct",
        "full_sharpe", "full_cagr_pct", "full_total_return_pct",
    ]
    print(valid_df[show_cols].head(min(8, len(valid_df))).to_string(index=False, float_format=lambda x: f"{x:.3f}"))

    b = valid_df.iloc[0]
    print("\n=== Recommended run command ===\n")
    cmd = (
        f"cd {_REPO} && PYTHONUNBUFFERED=1 .venv/bin/python "
        f"RenTech/strategy_stack/run_markov_vix_futures_backtest.py "
        f"--start {args.train_start} --fixed-expiry "
        f"--edge-threshold {b['edge_threshold']} --kelly-mult {b['kelly_mult']} "
        f"--max-short {b['max_short']} --gross-cap {b['gross_cap']} "
        f"--n-sims 3000 --matrix-refresh {args.matrix_refresh}"
    )
    if b["rel_spread"]:
        cmd += f" --rel-spread --rel-threshold {b['rel_threshold']}"
    if b["vix_regime"]:
        cmd += " --vix-regime"
    if not b["allow_short"]:
        cmd += " --long-only"
    print(cmd)


if __name__ == "__main__":
    main()
