#!/usr/bin/env python3
"""
Optimize vol-targeted French decile spreads (default 10% ann vol).

Sweeps momentum (2-12), variance shorts, short-term reversal, blends, and
tuning of lookback / target vol on top candidates vs baseline Hi PRIOR − Hi 10.

Example::

    cd /Users/robzingale/trading_bot
    PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/optimize_french_vol_target.py \\
        --start 1963-07-01 \\
        --out-prefix RenTech/data/logs/french_vol_target_opt
"""

from __future__ import annotations

import argparse
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.french_decile_loader import (
    load_momentum_deciles,
    load_short_term_reversal_deciles,
    load_variance_deciles,
    pct_to_decimal,
    vol_scale_monthly,
)

LOGS = _REPO / "RenTech" / "data" / "logs"

MOM_ORDER = [
    "Lo PRIOR", "PRIOR 2", "PRIOR 3", "PRIOR 4", "PRIOR 5",
    "PRIOR 6", "PRIOR 7", "PRIOR 8", "PRIOR 9", "Hi PRIOR",
]
VAR_ORDER = ["Lo 10", "Dec 2", "Dec 3", "Dec 4", "Dec 5", "Dec 6", "Dec 7", "Dec 8", "Dec 9", "Hi 10"]

BASELINE_RAW = ("Hi PRIOR", "Hi 10", "mom_212")
BASELINE_TARGET = 0.10
BASELINE_LOOKBACK = 36


def _align(frames: list[pd.DataFrame], start: str | None, end: str | None) -> list[pd.DataFrame]:
    idx = frames[0].index
    for f in frames[1:]:
        idx = idx.intersection(f.index)
    out = [f.loc[idx].copy() for f in frames]
    if start:
        ts = pd.Timestamp(start)
        out = [f.loc[f.index >= ts] for f in out]
    if end:
        te = pd.Timestamp(end)
        out = [f.loc[f.index <= te] for f in out]
    return out


def _stats(r: pd.Series, capital: float = 100_000.0) -> dict:
    r = r.dropna().astype(float)
    if len(r) < 24:
        return {"sharpe": np.nan, "cagr_pct": np.nan, "max_dd_pct": np.nan, "calmar": np.nan, "vol_ann_pct": np.nan}
    eq = capital * (1.0 + r).cumprod()
    dd = eq / eq.cummax() - 1.0
    sd = r.std(ddof=1)
    sharpe = float(r.mean() / sd * np.sqrt(12.0)) if sd > 1e-12 else np.nan
    years = len(r) / 12.0
    cagr = float((eq.iloc[-1] / capital) ** (1.0 / years) - 1.0) if years > 0 else np.nan
    max_dd = float(dd.min())
    calmar = float(cagr / abs(max_dd)) if max_dd < -1e-6 else np.nan
    return {
        "sharpe": round(sharpe, 4),
        "cagr_pct": round(100.0 * cagr, 4),
        "max_dd_pct": round(100.0 * max_dd, 4),
        "vol_ann_pct": round(100.0 * sd * np.sqrt(12.0), 4),
        "calmar": round(calmar, 4) if calmar == calmar else None,
        "ending_equity_usd": round(float(eq.iloc[-1]), 2),
    }


def _vt(
    raw: pd.Series,
    *,
    target_ann: float = BASELINE_TARGET,
    lookback: int = BASELINE_LOOKBACK,
    scale_cap: float = 2.0,
) -> pd.Series:
    scaled, _ = vol_scale_monthly(
        raw, target_ann=target_ann, lookback=lookback, scale_cap=scale_cap,
    )
    return scaled


def optimize(
    *,
    start: str | None,
    end: str | None,
    weighting: str,
    capital: float,
    out_prefix: Path,
    verbose: bool,
) -> pd.DataFrame:
    mom = load_momentum_deciles(weighting=weighting)  # type: ignore[arg-type]
    var = load_variance_deciles(weighting=weighting)  # type: ignore[arg-type]
    st = load_short_term_reversal_deciles(weighting=weighting)  # type: ignore[arg-type]
    mom, var, st = _align([mom, var, st], start, end)

    mom_r = {c: pct_to_decimal(mom[c]) for c in MOM_ORDER}
    var_r = {c: pct_to_decimal(var[c]) for c in VAR_ORDER}
    st_r = {c: pct_to_decimal(st[c]) for c in MOM_ORDER}

    lowvol = var_r["Lo 10"] - var_r["Hi 10"]
    baseline_raw = mom_r["Hi PRIOR"] - var_r["Hi 10"]

    rows: list[dict] = []

    def add(
        family: str,
        name: str,
        raw: pd.Series,
        *,
        target_ann: float = BASELINE_TARGET,
        lookback: int = BASELINE_LOOKBACK,
        scale_cap: float = 2.0,
        long_leg: str = "",
        short_leg: str = "",
        already_vol_scaled: bool = False,
    ) -> None:
        scaled = raw if already_vol_scaled else _vt(
            raw, target_ann=target_ann, lookback=lookback, scale_cap=scale_cap,
        )
        st_d = _stats(scaled, capital)
        base_st = _stats(_vt(baseline_raw), capital)
        rows.append({
            "family": family,
            "name": name,
            "long_leg": long_leg,
            "short_leg": short_leg,
            "target_vol_ann": target_ann,
            "vol_lookback": lookback,
            "scale_cap": scale_cap,
            "sharpe_vs_baseline": round(st_d["sharpe"] - base_st["sharpe"], 4)
            if st_d["sharpe"] == st_d["sharpe"] and base_st["sharpe"] == base_st["sharpe"]
            else None,
            "dd_vs_baseline_pp": round(st_d["max_dd_pct"] - base_st["max_dd_pct"], 2)
            if st_d["max_dd_pct"] == st_d["max_dd_pct"]
            else None,
            **st_d,
        })

    # Baseline
    add("baseline", "Hi PRIOR - Hi 10 @ 10%/36m", baseline_raw, long_leg="Hi PRIOR", short_leg="Hi 10")

    # --- Momentum (2-12) long x variance short ---
    for mi in range(5, 10):
        for vi in range(5, 10):
            ml, vs = MOM_ORDER[mi], VAR_ORDER[vi]
            add("mom_x_var", f"{ml} - {vs}", mom_r[ml] - var_r[vs], long_leg=ml, short_leg=vs)

    # --- Narrow momentum (2-12) with Hi 10 short ---
    for mi in range(6, 10):
        for lo_i in range(0, mi):
            ml, sl = MOM_ORDER[mi], MOM_ORDER[lo_i]
            add("mom_narrow", f"{ml} - {sl}", mom_r[ml] - mom_r[sl], long_leg=ml, short_leg=sl)

    # --- Short-term reversal (Lo - Hi on 1-month prior) vol-targeted ---
    st_raw = st_r["Lo PRIOR"] - st_r["Hi PRIOR"]
    add("st_rev", "ST rev Lo-Hi (1m)", st_raw, long_leg="ST Lo", short_leg="ST Hi")

    # --- Blends of raw spreads, then single vol target ---
    mom_hi = baseline_raw
    for w in (0.25, 0.5, 0.75):
        blend = w * mom_hi + (1.0 - w) * lowvol
        add("blend_mom_lowvol", f"{w:.0%} mom_hi_hi10 + {1-w:.0%} lowvol", blend)
        blend2 = w * mom_hi + (1.0 - w) * st_raw
        add("blend_mom_strev", f"{w:.0%} mom_hi_hi10 + {1-w:.0%} st_rev", blend2)

    # --- Vol-scale each leg separately, then combine (matched risk) ---
    vt_mom = _vt(mom_hi)
    vt_low = _vt(lowvol)
    vt_st = _vt(st_raw)
    for w in (0.5, 0.75):
        add(
            "blend_vt_legs", f"{w:.0%} vt_mom + {1-w:.0%} vt_lowvol",
            w * vt_mom + (1 - w) * vt_low, already_vol_scaled=True,
        )
        add(
            "blend_vt_legs", f"{w:.0%} vt_mom + {1-w:.0%} vt_strev",
            w * vt_mom + (1 - w) * vt_st, already_vol_scaled=True,
        )

    # --- Top raw candidates: tune lookback & target vol ---
    tune_candidates = [
        ("Hi PRIOR - Hi 10", baseline_raw),
        ("Hi PRIOR - Dec 9", mom_r["Hi PRIOR"] - var_r["Dec 9"]),
        ("PRIOR 9 - Hi 10", mom_r["PRIOR 9"] - var_r["Hi 10"]),
        ("Hi PRIOR - PRIOR 9", mom_r["Hi PRIOR"] - mom_r["PRIOR 9"]),
        ("75% vt_mom + 25% vt_low", 0.75 * vt_mom + 0.25 * vt_low),
    ]
    for label, raw in tune_candidates:
        if isinstance(raw, pd.Series) and "vt_" in label:
            # already scaled combo — only tune lookback on components not applicable; skip param grid
            continue
        for lb in (24, 36, 48):
            for tv in (0.08, 0.10, 0.12):
                add("tune", f"{label} @ {tv:.0%}/{lb}m", raw, target_ann=tv, lookback=lb)

    df = pd.DataFrame(rows)
    out_prefix.parent.mkdir(parents=True, exist_ok=True)
    df.to_csv(f"{out_prefix}_grid.csv", index=False)

    baseline_row = df[df["family"] == "baseline"].iloc[0]
    better_sharpe = df[
        (df["sharpe"] > baseline_row["sharpe"])
        & (df["max_dd_pct"] >= baseline_row["max_dd_pct"] - 2.0)
    ].sort_values(["sharpe", "max_dd_pct"], ascending=[False, False])

    better_dd = df[
        (df["max_dd_pct"] > baseline_row["max_dd_pct"])
        & (df["sharpe"] >= baseline_row["sharpe"] - 0.02)
    ].sort_values(["max_dd_pct", "sharpe"], ascending=[False, False])

    best_calmar = df.sort_values("calmar", ascending=False).head(20)

    cmd = (
        f"cd {_REPO} && PYTHONUNBUFFERED=1 .venv/bin/python "
        f"RenTech/strategy_stack/optimize_french_vol_target.py "
        f"--weighting {weighting} --out-prefix {out_prefix}"
    )
    if start:
        cmd += f" --start {start}"

    summary = {
        "window": {"start": str(mom.index.min().date()), "end": str(mom.index.max().date())},
        "weighting": weighting,
        "baseline": baseline_row.to_dict(),
        "n_candidates": len(df),
        "beat_baseline_sharpe_similar_dd": better_sharpe.head(10).to_dict(orient="records"),
        "beat_baseline_dd_similar_sharpe": better_dd.head(10).to_dict(orient="records"),
        "top_calmar": best_calmar.head(10).to_dict(orient="records"),
        "recommended": None,
    }

    # Pick recommendation: max Sharpe among DD within 3pp of baseline (less negative = better)
    dd_floor = baseline_row["max_dd_pct"] - 3.0
    rec_pool = df[(df["max_dd_pct"] >= dd_floor) & (df["sharpe"] >= baseline_row["sharpe"])]
    if not rec_pool.empty:
        rec = rec_pool.sort_values(["sharpe", "calmar"], ascending=False).iloc[0]
        summary["recommended"] = rec.to_dict()

    with open(f"{out_prefix}_summary.json", "w") as f:
        json.dump(summary, f, indent=2, default=str)

    lines = [
        "=== French vol-target optimization ===",
        f"Window: {summary['window']['start']} → {summary['window']['end']}  |  {weighting}-weight",
        "",
        "Baseline (Hi PRIOR − Hi 10 @ 10% vol, 36m):",
        f"  Sharpe {baseline_row['sharpe']:.3f}  CAGR {baseline_row['cagr_pct']:.2f}%  "
        f"MaxDD {baseline_row['max_dd_pct']:.2f}%  Calmar {baseline_row['calmar']}",
        "",
    ]
    if summary["recommended"]:
        r = summary["recommended"]
        lines.append("Recommended upgrade:")
        lines.append(
            f"  [{r['family']}] {r['name']}: Sharpe {r['sharpe']:.3f}  CAGR {r['cagr_pct']:.2f}%  "
            f"MaxDD {r['max_dd_pct']:.2f}%  (ΔSharpe {r.get('sharpe_vs_baseline')}, "
            f"ΔDD {r.get('dd_vs_baseline_pp')} pp)"
        )
    else:
        lines.append("No candidate beat baseline on Sharpe with similar DD; see top Calmar below.")

    lines.extend(["", "=== Top 10 Calmar ==="])
    for _, r in best_calmar.head(10).iterrows():
        lines.append(
            f"  [{r['family']}] {r['name']}: Sharpe {r['sharpe']:.3f}  "
            f"CAGR {r['cagr_pct']:.1f}%  MaxDD {r['max_dd_pct']:.1f}%  Calmar {r['calmar']}"
        )

    lines.extend(["", "=== Better DD (Sharpe ≥ baseline − 0.02) ==="])
    for _, r in better_dd.head(8).iterrows():
        lines.append(
            f"  [{r['family']}] {r['name']}: Sharpe {r['sharpe']:.3f}  MaxDD {r['max_dd_pct']:.1f}%"
        )

    lines.extend(["", "Command:", f"  {cmd}"])
    report = "\n".join(lines) + "\n"
    with open(f"{out_prefix}_report.txt", "w") as f:
        f.write(report)
    if verbose:
        print(report)

    # Export best strategy monthly if recommended differs from baseline
    if summary["recommended"] and summary["recommended"]["name"] != baseline_row["name"]:
        best_name = summary["recommended"]["name"]
        # Rebuild raw for recommended from row metadata when possible
        rec_row = summary["recommended"]
        if rec_row.get("long_leg") and rec_row.get("short_leg"):
            ll, sl = rec_row["long_leg"], rec_row["short_leg"]
            if ll in mom_r and sl in var_r:
                raw_best = mom_r[ll] - var_r[sl]
            elif ll in mom_r and sl in mom_r:
                raw_best = mom_r[ll] - mom_r[sl]
            else:
                raw_best = baseline_raw
        else:
            raw_best = baseline_raw
        tv = float(rec_row.get("target_vol_ann", BASELINE_TARGET))
        lb = int(rec_row.get("vol_lookback", BASELINE_LOOKBACK))
        scaled, scales = vol_scale_monthly(raw_best, target_ann=tv, lookback=lb)
        eq = capital * (1.0 + scaled).cumprod()
        pd.DataFrame({
            "date": scaled.index.strftime("%Y-%m-%d"),
            "monthly_return": scaled.values,
            "equity_usd": eq.values,
            "vol_scale": scales.values,
            "raw_spread_return": raw_best.reindex(scaled.index).values,
            "strategy_name": best_name,
        }).to_csv(f"{out_prefix}_recommended_monthly.csv", index=False)

    return df


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__)
    ap.add_argument("--start", default="1963-07-01")
    ap.add_argument("--end", default=None)
    ap.add_argument("--weighting", choices=("value", "equal"), default="value")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--out-prefix", type=Path, default=LOGS / "french_vol_target_opt")
    ap.add_argument("-q", "--quiet", action="store_true")
    args = ap.parse_args()
    optimize(
        start=args.start,
        end=args.end,
        weighting=args.weighting,
        capital=args.capital,
        out_prefix=args.out_prefix,
        verbose=not args.quiet,
    )


if __name__ == "__main__":
    main()
