#!/usr/bin/env python3
"""
Optimize VRP sleeve risk fractions (R1/R2a/R2b/R3/R4) for Sharpe.

Sleeve mapping:
  - R1  -> pmcc
  - R2a -> diagonal
  - R2b -> r2_spread
  - R3  -> naked
  - R4  -> credit_spread

Objective uses a realized daily equity curve (cash equity from closed trades, forward-filled
over trading days) and annualized daily Sharpe (sqrt(252)).
"""

from __future__ import annotations

import argparse
import json
import math
import random
import sys
from pathlib import Path
from typing import Any, cast

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.theta_chunks_loader import ThetaChunksLoader, theta_chunks_date_bounds
from RenTech.strategy_stack.vrp_backtest_theta import (
    load_spy_vix_from_yfinance,
    normalize_spy_df,
    trading_days_intersecting_spy,
)
from RenTech.strategy_stack.vrp_backtester import VRPBacktester
from RenTech.strategy_stack.vrp_strategy_config import (
    DEFAULT_STRATEGY_CONFIG_PATH,
    apply_strategy_params_to_vrp_backtester_module,
    load_strategy_config_file,
)

_DEFAULT_THETA_DIR = _REPO_ROOT / "RenTech" / "data" / "theta_chunks"


def _daily_realized_equity(bt: VRPBacktester, days: list[pd.Timestamp]) -> pd.Series:
    """Forward-fill closed-trade equity over the trading-day index."""
    eq_points = {pd.Timestamp(d).normalize(): float(v) for d, v in bt._equity_curve}  # noqa: SLF001
    idx = pd.DatetimeIndex([pd.Timestamp(d).normalize() for d in days], name="date")
    s = pd.Series(index=idx, dtype=float)
    for d in idx:
        if d in eq_points:
            s.loc[d] = eq_points[d]
    s = s.ffill()
    if s.isna().all():
        s[:] = float(bt.initial_capital)
    else:
        s = s.fillna(float(bt.initial_capital))
    return s


def _sharpe_daily(realized_eq: pd.Series) -> float:
    r = realized_eq.pct_change().fillna(0.0)
    sd = float(r.std(ddof=1))
    if sd <= 1e-12:
        return float("nan")
    mu = float(r.mean())
    return (mu / sd) * math.sqrt(252.0)


def _max_dd(realized_eq: pd.Series) -> float:
    peak = realized_eq.cummax()
    dd = (peak - realized_eq) / peak.where(peak > 0, 1.0)
    return float(dd.max())


def _sample_between(lo: float, hi: float) -> float:
    return random.random() * (hi - lo) + lo


def main() -> None:
    ap = argparse.ArgumentParser(description="Optimize VRP sleeve fractions for Sharpe (random search).")
    ap.add_argument("--theta-dir", type=Path, default=_DEFAULT_THETA_DIR)
    ap.add_argument("--capital", type=float, default=100_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("--strategy-config", type=Path, default=None)
    ap.add_argument("--max-evals", type=int, default=12)
    ap.add_argument("--seed", type=int, default=7)
    ap.add_argument("--top-k", type=int, default=5)
    ap.add_argument("--max-dd", type=float, default=0.0, help="Optional max drawdown cap (fraction). 0 disables.")
    ap.add_argument("--no-vol-scaling", action="store_true")
    ap.add_argument("--no-r2-crossover-filters", action="store_true")
    ap.add_argument(
        "--bounds-json",
        type=str,
        default="",
        help='JSON object with min/max bounds by sleeve key, e.g. {"pmcc":[0.01,0.04],...}',
    )
    ap.add_argument(
        "--out-json",
        type=Path,
        default=_REPO_ROOT / "RenTech" / "strategy_stack" / "vrp_sleeve_sharpe_optimized.json",
    )
    args = ap.parse_args()

    random.seed(int(args.seed))
    cfg_path = args.strategy_config.expanduser() if args.strategy_config is not None else DEFAULT_STRATEGY_CONFIG_PATH
    if not cfg_path.is_file():
        raise FileNotFoundError(f"strategy config not found: {cfg_path}")
    cfg = load_strategy_config_file(cfg_path)
    apply_strategy_params_to_vrp_backtester_module(cfg.strategy_params)

    theta_dir = args.theta_dir.expanduser()
    if not theta_dir.is_dir():
        raise NotADirectoryError(f"--theta-dir is not a directory: {theta_dir}")

    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 args.start.strip():
        t0 = pd.Timestamp(args.start.strip())
        days = [d for d in days if d >= t0]
    if args.end.strip():
        t1 = pd.Timestamp(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 RuntimeError("no trading days after filters")

    bounds: dict[str, list[float]] = {
        "pmcc": [0.005, 0.050],
        "diagonal": [0.020, 0.120],
        "r2_spread": [0.005, 0.040],
        "naked": [0.005, 0.050],
        "credit_spread": [0.005, 0.050],
    }
    if args.bounds_json.strip():
        usr = json.loads(args.bounds_json)
        for k, v in usr.items():
            if k in bounds and isinstance(v, list) and len(v) == 2:
                bounds[k] = [float(v[0]), float(v[1])]

    bt_base: dict[str, Any] = {
        "initial_capital": float(args.capital),
        "spy_df": spy_wide,
        "vol_risk_scaling": not bool(args.no_vol_scaling),
        "r2_crossover_filters": not bool(args.no_r2_crossover_filters),
        "dd_risk_scaling": False,
        "overlay_risk_fractions": dict(cfg.overlay_risk_fractions) if cfg.overlay_risk_fractions else None,
        "overlay_risk_cap_frac": cfg.overlay_risk_cap_frac,
        # Turn off hard cap during search so high-R2 candidates are testable.
        "total_risk_cap_frac": None,
    }

    base = dict(cfg.sleeve_risk_fractions)
    candidates: list[dict[str, float]] = [cast(dict[str, float], base)]
    for _ in range(max(0, int(args.max_evals) - 1)):
        c = dict(base)
        for k in ("pmcc", "diagonal", "r2_spread", "naked", "credit_spread"):
            lo, hi = bounds[k]
            c[k] = _sample_between(lo, hi)
        candidates.append(cast(dict[str, float], c))

    rows: list[dict[str, Any]] = []
    for i, sleeves in enumerate(candidates, start=1):
        bt = VRPBacktester(ld, **bt_base, sleeve_risk_fractions=cast(dict, sleeves))
        bt.run_backtest(trading_days=days, show_progress=False)
        m = bt.metrics()
        eq = _daily_realized_equity(bt, days)
        sh = _sharpe_daily(eq)
        dd = _max_dd(eq)
        feasible = (float(args.max_dd) <= 0.0) or (dd <= float(args.max_dd))
        row = {
            "rank_hint": i,
            "pmcc": sleeves["pmcc"],
            "diagonal": sleeves["diagonal"],
            "r2_spread": sleeves["r2_spread"],
            "naked": sleeves["naked"],
            "credit_spread": sleeves["credit_spread"],
            "sharpe_daily_realized": sh,
            "max_dd_realized": dd,
            "engine_total_return": float(m["total_return"]),
            "engine_cagr": float(m["cagr"]) if isinstance(m["cagr"], float) and math.isfinite(float(m["cagr"])) else None,
            "ending_capital": float(m["ending_capital"]),
            "total_trades": int(m["total_trades"]),
            "feasible_dd": bool(feasible),
        }
        rows.append(row)
        shs = f"{sh:.3f}" if isinstance(sh, float) and math.isfinite(sh) else "nan"
        print(
            f"[{i}/{len(candidates)}] sharpe={shs} dd={dd:.3%} "
            f"R1={sleeves['pmcc']:.2%} R2a={sleeves['diagonal']:.2%} R2b={sleeves['r2_spread']:.2%} "
            f"R3={sleeves['naked']:.2%} R4={sleeves['credit_spread']:.2%}",
            flush=True,
        )

    valid = [r for r in rows if r["feasible_dd"] and isinstance(r["sharpe_daily_realized"], float)]
    valid = [r for r in valid if math.isfinite(float(r["sharpe_daily_realized"]))]
    ranked = sorted(valid, key=lambda x: float(x["sharpe_daily_realized"]), reverse=True)
    top_k = ranked[: max(1, int(args.top_k))]
    best = top_k[0] if top_k else None

    out = {
        "window": {
            "first_day": str(days[0].date()),
            "last_day": str(days[-1].date()),
            "n_days": len(days),
        },
        "capital": float(args.capital),
        "max_evals": int(args.max_evals),
        "max_dd_constraint": float(args.max_dd),
        "best": best,
        "top_k": top_k,
        "all_trials": rows,
    }
    args.out_json.parent.mkdir(parents=True, exist_ok=True)
    args.out_json.write_text(json.dumps(out, indent=2), encoding="utf-8")
    print(f"\nWrote optimization report: {args.out_json}")
    if best:
        print(
            "Best sleeves: "
            f"R1={best['pmcc']:.2%}, R2a={best['diagonal']:.2%}, R2b={best['r2_spread']:.2%}, "
            f"R3={best['naked']:.2%}, R4={best['credit_spread']:.2%} | "
            f"Sharpe={best['sharpe_daily_realized']:.3f}, DD={best['max_dd_realized']:.2%}"
        )
    else:
        print("No feasible candidate met constraints.")


if __name__ == "__main__":
    main()

