#!/usr/bin/env python3
"""
Find minimum max-DD French vol-target variants (monthly portfolio returns).

Example::

    cd /Users/robzingale/trading_bot
    PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/optimize_french_low_dd.py \\
        --start 1963-07-01 --min-sharpe 0.35 \\
        --out-prefix RenTech/data/logs/french_low_dd_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"


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:
        out = [f.loc[f.index >= pd.Timestamp(start)] for f in out]
    if end:
        out = [f.loc[f.index <= pd.Timestamp(end)] 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}
    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)
    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,
    }


def _dd_overlay(r: pd.Series, *, trigger: float = -0.10, exposure: float = 0.5) -> pd.Series:
    """Cut exposure to ``exposure`` when strategy equity is in drawdown below ``trigger``."""
    r = r.astype(float).copy()
    eq = (1.0 + r).cumprod()
    peak = eq.cummax()
    dd = eq / peak - 1.0
    mult = pd.Series(1.0, index=r.index)
    mult.loc[dd < trigger] = exposure
    return r * mult


def run(
    *,
    start: str | None,
    end: str | None,
    weighting: str,
    min_sharpe: float,
    min_cagr_pct: float,
    capital: float,
    out_prefix: Path,
    verbose: bool,
) -> None:
    mom, var, st = _align(
        [
            load_momentum_deciles(weighting=weighting),  # type: ignore[arg-type]
            load_variance_deciles(weighting=weighting),  # type: ignore[arg-type]
            load_short_term_reversal_deciles(weighting=weighting),  # type: ignore[arg-type]
        ],
        start,
        end,
    )
    mom_hi = pct_to_decimal(mom["Hi PRIOR"]) - pct_to_decimal(var["Hi 10"])
    st_rev = pct_to_decimal(st["Lo PRIOR"]) - pct_to_decimal(st["Hi PRIOR"])
    lowvol = pct_to_decimal(var["Lo 10"]) - pct_to_decimal(var["Hi 10"])

    rows: list[dict] = []

    def emit(name: str, series: pd.Series, family: str, **kw) -> None:
        st = _stats(series, capital)
        if st["sharpe"] != st["sharpe"] or st["sharpe"] < min_sharpe:
            return
        if st["cagr_pct"] < min_cagr_pct:
            return
        rows.append({"family": family, "name": name, **kw, **st})

    # --- Vol target grids: mom spread, ST rev, lowvol, blends ---
    for tv in (0.04, 0.05, 0.06, 0.08, 0.10):
        for lb in (24, 36, 48):
            for cap in (1.0, 1.5):
                vt_m, _ = vol_scale_monthly(mom_hi, target_ann=tv, lookback=lb, scale_cap=cap)
                vt_s, _ = vol_scale_monthly(st_rev, target_ann=tv, lookback=lb, scale_cap=cap)
                vt_l, _ = vol_scale_monthly(lowvol, target_ann=tv, lookback=lb, scale_cap=cap)
                emit(
                    f"mom_hi_hi10 @ {tv:.0%} vol {lb}m",
                    vt_m,
                    "vol_mom",
                    target_vol=tv,
                    lookback=lb,
                    scale_cap=cap,
                )
                emit(f"st_rev @ {tv:.0%} vol {lb}m", vt_s, "vol_st", target_vol=tv, lookback=lb)
                emit(f"lowvol @ {tv:.0%} vol {lb}m", vt_l, "vol_lowvol", target_vol=tv, lookback=lb)
                for wm in (0.0, 0.25, 0.5, 0.75, 1.0):
                    w = wm
                    blend = w * vt_m + (1.0 - w) * vt_s
                    emit(
                        f"{w:.0%} mom + {1-w:.0%} st @ {tv:.0%}/{lb}m",
                        blend,
                        "blend_mom_st",
                        target_vol=tv,
                        lookback=lb,
                        mom_weight=wm,
                    )
                for wm in (0.5, 0.75):
                    blend2 = wm * vt_m + (1.0 - wm) * vt_l
                    emit(
                        f"{wm:.0%} mom + {1-wm:.0%} lowvol @ {tv:.0%}/{lb}m",
                        blend2,
                        "blend_mom_lowvol",
                        target_vol=tv,
                        lookback=lb,
                    )

    # --- Fixed exposure fractions on vol-targeted 10%/36m mom ---
    vt_m10, _ = vol_scale_monthly(mom_hi, target_ann=0.10, lookback=36)
    vt_s10, _ = vol_scale_monthly(st_rev, target_ann=0.10, lookback=36)
    for frac in (0.15, 0.20, 0.25, 0.35, 0.50):
        emit(f"{frac:.0%} x vt_mom_10", vt_m10 * frac, "leverage", fraction=frac)
        emit(
            f"{frac:.0%} x (75% mom + 25% st)",
            (0.75 * vt_m10 + 0.25 * vt_s10) * frac,
            "leverage_blend",
            fraction=frac,
        )

    # --- DD overlay on best structural candidates ---
    candidates = {
        "blend_50_st_10_36": 0.5 * vt_m10 + 0.5 * vt_s10,
        "blend_75_st_10_36": 0.75 * vt_m10 + 0.25 * vt_s10,
        "mom_8_24": vol_scale_monthly(mom_hi, target_ann=0.08, lookback=24)[0],
        "blend_50_st_8_24": (
            0.5 * vol_scale_monthly(mom_hi, target_ann=0.08, lookback=24)[0]
            + 0.5 * vol_scale_monthly(st_rev, target_ann=0.08, lookback=24)[0]
        ),
        "blend_25_st_6_36": (
            0.25 * vol_scale_monthly(mom_hi, target_ann=0.06, lookback=36)[0]
            + 0.75 * vol_scale_monthly(st_rev, target_ann=0.06, lookback=36)[0]
        ),
    }
    for label, ser in candidates.items():
        for trig in (-0.08, -0.10, -0.12):
            for exp in (0.25, 0.5, 0.0):
                emit(
                    f"{label} DD>{abs(trig)*100:.0f}%→{exp:.0%}expo",
                    _dd_overlay(ser, trigger=trig, exposure=exp),
                    "dd_overlay",
                    base=label,
                    trigger=trig,
                    cut_exposure=exp,
                )

    df = pd.DataFrame(rows)
    if df.empty:
        raise RuntimeError("No strategies passed filters; relax --min-sharpe or --min-cagr-pct")

    out_prefix.parent.mkdir(parents=True, exist_ok=True)
    df.sort_values("max_dd_pct", ascending=False).to_csv(f"{out_prefix}_all.csv", index=False)

    by_dd = df.sort_values("max_dd_pct", ascending=False)
    by_calmar = df.sort_values("calmar", ascending=False)

    # Picks
    best_dd = by_dd.iloc[0]
    best_dd_sharpe = by_dd.sort_values("sharpe", ascending=False).iloc[0]
    best_calmar_under_15 = df[df["max_dd_pct"] > -15].sort_values("calmar", ascending=False)
    best_calmar_under_20 = df[df["max_dd_pct"] > -20].sort_values("calmar", ascending=False)

    cmd = (
        f"cd {_REPO} && PYTHONUNBUFFERED=1 .venv/bin/python "
        f"RenTech/strategy_stack/optimize_french_low_dd.py "
        f"--out-prefix {out_prefix} --min-sharpe {min_sharpe} --min-cagr-pct {min_cagr_pct}"
    )
    if start:
        cmd += f" --start {start}"

    summary = {
        "window": {"start": str(mom.index.min().date()), "end": str(mom.index.max().date())},
        "filters": {"min_sharpe": min_sharpe, "min_cagr_pct": min_cagr_pct},
        "n_passing": len(df),
        "best_max_dd": best_dd.to_dict(),
        "best_max_dd_among_high_sharpe": best_dd_sharpe.to_dict(),
        "best_calmar_dd_better_than_-15": (
            best_calmar_under_15.head(5).to_dict(orient="records") if len(best_calmar_under_15) else []
        ),
        "best_calmar_dd_better_than_-20": best_calmar_under_20.head(8).to_dict(orient="records"),
        "command": cmd,
    }
    with open(f"{out_prefix}_summary.json", "w") as f:
        json.dump(summary, f, indent=2, default=str)

    lines = [
        "=== French low max-DD search ===",
        f"Window: {summary['window']['start']} → {summary['window']['end']}",
        f"Candidates passing Sharpe>={min_sharpe}, CAGR>={min_cagr_pct}%: {len(df)}",
        "",
        "Lowest max DD:",
        f"  [{best_dd['family']}] {best_dd['name']}",
        f"  Sharpe {best_dd['sharpe']:.3f}  CAGR {best_dd['cagr_pct']:.2f}%  MaxDD {best_dd['max_dd_pct']:.2f}%",
        "",
        "Best Sharpe among low-DD set:",
        f"  [{best_dd_sharpe['family']}] {best_dd_sharpe['name']}",
        f"  Sharpe {best_dd_sharpe['sharpe']:.3f}  CAGR {best_dd_sharpe['cagr_pct']:.2f}%  MaxDD {best_dd_sharpe['max_dd_pct']:.2f}%",
        "",
        "=== Top 12 by max DD (least negative) ===",
    ]
    for _, r in by_dd.head(12).iterrows():
        lines.append(
            f"  {r['max_dd_pct']:>7.2f}% DD  Sharpe {r['sharpe']:.3f}  CAGR {r['cagr_pct']:>6.2f}%  | {r['name']}"
        )
    lines.extend(["", "=== Best Calmar with DD better than −20% ==="])
    for _, r in best_calmar_under_20.head(8).iterrows():
        lines.append(
            f"  {r['max_dd_pct']:>7.2f}% DD  Calmar {r['calmar']:.3f}  Sharpe {r['sharpe']:.3f}  | {r['name']}"
        )
    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)


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("--min-sharpe", type=float, default=0.35)
    ap.add_argument("--min-cagr-pct", type=float, default=1.0, help="Minimum CAGR %% to keep")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--out-prefix", type=Path, default=LOGS / "french_low_dd_opt")
    ap.add_argument("-q", "--quiet", action="store_true")
    args = ap.parse_args()
    run(
        start=args.start,
        end=args.end,
        weighting=args.weighting,
        min_sharpe=args.min_sharpe,
        min_cagr_pct=args.min_cagr_pct,
        capital=args.capital,
        out_prefix=args.out_prefix,
        verbose=not args.quiet,
    )


if __name__ == "__main__":
    main()
