#!/usr/bin/env python3
"""
Sweep Ken French momentum + variance decile portfolios for L/S and risk-managed variants.

Goal: find high Sharpe / CAGR strategies with **lower max drawdown** than classic momentum.

Example::

    cd /Users/robzingale/trading_bot
    PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/discover_french_decile_strategies.py \\
        --start 1963-07-01 --min-sharpe 0.3 \\
        --out-prefix RenTech/data/logs/french_decile_discovery
"""

from __future__ import annotations

import argparse
import json
import sys
from itertools import combinations
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_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"]
VAR_Q_ORDER = ["Lo 20", "Qnt 2", "Qnt 3", "Qnt 4", "Hi 20"]


def _load_var_quintiles(weighting: str) -> pd.DataFrame:
    from RenTech.strategy_stack.french_decile_loader import FRENCH_DIR

    path = FRENCH_DIR / "Portfolios_Formed_on_VAR.csv"
    text = path.read_text(encoding="utf-8", errors="replace")
    lines = text.splitlines()
    title = (
        "  Value Weighted Returns -- Monthly"
        if weighting == "value"
        else "  Equal Weighted Returns -- Monthly"
    )
    from RenTech.strategy_stack.french_decile_loader import _section_monthly_table

    df = _section_monthly_table(lines, title)
    if df is None:
        raise ValueError("variance table missing")
    cols = [c for c in VAR_Q_ORDER if c in df.columns]
    return df[cols]


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 {"n_months": len(r), "sharpe": np.nan, "max_dd_pct": np.nan, "cagr_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) 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 {
        "n_months": int(len(r)),
        "sharpe": round(sharpe, 4),
        "cagr_pct": round(100.0 * cagr, 4),
        "total_return_pct": round(100.0 * (eq.iloc[-1] / capital - 1.0), 2),
        "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 discover(
    *,
    start: str | None,
    end: str | None,
    weighting: str,
    min_sharpe: float,
    capital: float,
    out_prefix: Path,
    verbose: bool,
) -> pd.DataFrame:
    mom = load_momentum_deciles(weighting=weighting)  # type: ignore[arg-type]
    var_d = load_variance_deciles(weighting=weighting)  # type: ignore[arg-type]
    var_q = _load_var_quintiles(weighting)
    mom, var_d, var_q = _align([mom, var_d, var_q], start, end)

    mom_r = {c: pct_to_decimal(mom[c]) for c in MOM_ORDER if c in mom.columns}
    var_r = {c: pct_to_decimal(var_d[c]) for c in VAR_ORDER}
    varq_r = {c: pct_to_decimal(var_q[c]) for c in VAR_Q_ORDER}

    rows: list[dict] = []

    def add_row(family: str, name: str, series: pd.Series, extra: dict | None = None) -> None:
        st = _stats(series, capital)
        if st.get("sharpe") is None or (st["sharpe"] != st["sharpe"]) or st["sharpe"] < min_sharpe:
            return
        row = {"family": family, "name": name, **st}
        if extra:
            row.update(extra)
        rows.append(row)

    # --- Within momentum (long higher decile, short lower) ---
    for i, hi in enumerate(MOM_ORDER):
        for lo in MOM_ORDER[:i]:
            add_row("mom_ls", f"{hi} - {lo}", mom_r[hi] - mom_r[lo])

    # --- Within variance deciles ---
    for i, lo in enumerate(VAR_ORDER):
        for hi in VAR_ORDER[i + 1 :]:
            add_row("var_ls", f"{lo} - {hi}", var_r[lo] - var_r[hi])

    # --- Variance quintiles ---
    for i, lo in enumerate(VAR_Q_ORDER):
        for hi in VAR_Q_ORDER[i + 1 :]:
            add_row("var_q_ls", f"{lo} - {hi}", varq_r[lo] - varq_r[hi])

    # --- Cross: mom long vs var short ---
    for mh in MOM_ORDER[5:]:  # top half mom
        for vs in VAR_ORDER[5:]:  # top half vol for short
            add_row("mom_x_var", f"{mh} - {vs}", mom_r[mh] - var_r[vs])
    for mh in MOM_ORDER[7:]:  # top 3 mom
        for vs in VAR_ORDER:
            add_row("mom_x_var_all", f"{mh} - {vs}", mom_r[mh] - var_r[vs])

    # --- Blends of classic legs ---
    mom_classic = mom_r["Hi PRIOR"] - mom_r["Lo PRIOR"]
    lowvol = var_r["Lo 10"] - var_r["Hi 10"]
    mom_hi_short_hi = mom_r["Hi PRIOR"] - var_r["Hi 10"]
    for w in (0.25, 0.5, 0.75):
        add_row("blend", f"{w:.0%} mom_ls + {1-w:.0%} lowvol", w * mom_classic + (1 - w) * lowvol)
        add_row("blend", f"{w:.0%} mom_hi_short_hi + {1-w:.0%} lowvol", w * mom_hi_short_hi + (1 - w) * lowvol)

    # --- Leverage fractions ---
    for frac in (0.25, 0.5, 0.75):
        add_row("leverage", f"{frac:.0%} x mom_ls", frac * mom_classic)
        add_row("leverage", f"{frac:.0%} x mom_hi_short_hi", frac * mom_hi_short_hi)
        add_row("leverage", f"{frac:.0%} x lowvol", frac * lowvol)

    # --- Vol targeting on key candidates ---
    for base_name, base_r in [
        ("mom_ls", mom_classic),
        ("mom_hi_short_hi", mom_hi_short_hi),
        ("lowvol", lowvol),
        ("PRIOR9-Lo", mom_r["PRIOR 9"] - mom_r["Lo PRIOR"]),
        ("Hi-Dec8", mom_r["Hi PRIOR"] - var_r["Dec 8"]),
        ("Hi-Dec9", mom_r["Hi PRIOR"] - var_r["Dec 9"]),
        ("PRIOR8-PRIOR2", mom_r["PRIOR 8"] - mom_r["PRIOR 2"]),
        ("Lo20-Hi20_q", varq_r["Lo 20"] - varq_r["Hi 20"]),
    ]:
        for tgt in (0.08, 0.10, 0.12):
            scaled, _ = vol_scale_monthly(base_r, target_ann=tgt)
            add_row(
                "vol_target",
                f"{base_name} @ {tgt:.0%} ann",
                scaled,
                {"base": base_name, "target_vol": tgt},
            )

    # --- Long-only (for DD comparison) ---
    if "Hi PRIOR" in mom_r:
        add_row("long_only", "Hi PRIOR", mom_r["Hi PRIOR"])
    if "Lo 10" in var_r:
        add_row("long_only", "Lo 10", var_r["Lo 10"])

    df = pd.DataFrame(rows)
    if df.empty:
        raise RuntimeError("No strategies passed min_sharpe filter")

    baseline_mask = (df["family"] == "mom_ls") & (df["name"] == "Hi PRIOR - Lo PRIOR")
    if baseline_mask.any():
        mom_dd = float(df.loc[baseline_mask, "max_dd_pct"].iloc[0])
        df["dd_vs_classic_mom_pct"] = df["max_dd_pct"] - mom_dd
    else:
        mom_dd = None
        df["dd_vs_classic_mom_pct"] = np.nan

    out_prefix.parent.mkdir(parents=True, exist_ok=True)
    df.sort_values(["max_dd_pct", "sharpe"], ascending=[False, False]).to_csv(
        f"{out_prefix}_all_ranked_by_dd.csv", index=False
    )
    df.sort_values("calmar", ascending=False).to_csv(f"{out_prefix}_all_ranked_by_calmar.csv", index=False)

    # Best under DD thresholds with min CAGR
    picks = []
    for dd_floor in (-70, -60, -50, -40, -35, -30, -25, -20):
        sub = df[(df["max_dd_pct"] >= dd_floor) & (df["cagr_pct"] > 0)].copy()
        if sub.empty:
            continue
        best = sub.sort_values(["sharpe", "cagr_pct"], ascending=False).head(15)
        best["dd_bucket"] = dd_floor
        picks.append(best)
    picks_df = pd.concat(picks, ignore_index=True)
    picks_df.to_csv(f"{out_prefix}_best_by_dd_bucket.csv", index=False)

    # Pareto: not dominated on (sharpe, max_dd) — higher sharpe AND higher (less negative) dd
    pareto_rows = []
    for _, row in df.iterrows():
        dominated = False
        for _, other in df.iterrows():
            if other.name == row.name and other.family == row.family:
                continue
            if other.sharpe >= row.sharpe and other.max_dd_pct >= row.max_dd_pct:
                if other.sharpe > row.sharpe or other.max_dd_pct > row.max_dd_pct:
                    dominated = True
                    break
        if not dominated:
            pareto_rows.append(row)
    pareto_df = pd.DataFrame(pareto_rows).sort_values("max_dd_pct", ascending=False)
    pareto_df.to_csv(f"{out_prefix}_pareto_sharpe_dd.csv", index=False)

    baseline = df[baseline_mask].iloc[0] if baseline_mask.any() else None
    hi_hi_mask = df["name"] == "Hi PRIOR - Hi 10"
    hi_hi = df[hi_hi_mask].iloc[0] if hi_hi_mask.any() else None

    top_dd = df[df["cagr_pct"] >= 4.0].sort_values("max_dd_pct", ascending=False).head(20)
    top_calmar = df.sort_values("calmar", ascending=False).head(20)

    summary = {
        "window": {"start": str(mom.index.min().date()), "end": str(mom.index.max().date())},
        "weighting": weighting,
        "n_strategies_tested": int(len(df)),
        "classic_mom_ls": baseline.to_dict() if baseline is not None else None,
        "mom_hi_short_hi_vol": hi_hi.to_dict() if hi_hi is not None else None,
        "best_cagr_with_dd_better_than_-40": top_dd[top_dd["max_dd_pct"] > -40].head(5).to_dict(orient="records"),
        "best_calmar_top5": top_calmar.head(5).to_dict(orient="records"),
        "best_dd_under_-25_sharpe": (
            df[(df["max_dd_pct"] > -25) & (df["cagr_pct"] > 2)]
            .sort_values("sharpe", ascending=False)
            .head(10)
            .to_dict(orient="records")
        ),
        "pareto_count": len(pareto_df),
    }

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

    lines = [
        "=== French decile strategy discovery ===",
        f"Window: {summary['window']['start']} → {summary['window']['end']}  |  weighting={weighting}",
        f"Strategies passing Sharpe >= {min_sharpe}: {len(df)}",
        "",
        "Baseline (classic momentum L/S):",
    ]
    if baseline is not None:
        lines.append(
            f"  Sharpe {baseline['sharpe']:.3f}  CAGR {baseline['cagr_pct']:.2f}%  "
            f"MaxDD {baseline['max_dd_pct']:.2f}%"
        )
    else:
        lines.append("  (not in filtered set)")
    if hi_hi is not None:
        lines.append(
            f"Prior best (Hi PRIOR - Hi 10): Sharpe {hi_hi['sharpe']:.3f}  "
            f"CAGR {hi_hi['cagr_pct']:.2f}%  MaxDD {hi_hi['max_dd_pct']:.2f}%"
        )
    lines.extend(["", "=== Top 12 by Calmar (CAGR / |MaxDD|) ==="])
    for _, r in top_calmar.head(12).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(["", "=== Best MaxDD among CAGR >= 4% ==="])
    for _, r in top_dd.head(12).iterrows():
        lines.append(
            f"  [{r['family']}] {r['name']}: Sharpe {r['sharpe']:.3f}  "
            f"CAGR {r['cagr_pct']:.1f}%  MaxDD {r['max_dd_pct']:.1f}%"
        )
    lines.extend(["", "=== Lowest DD with Sharpe >= 0.5 and CAGR >= 2% ==="])
    low = df[(df["sharpe"] >= 0.5) & (df["cagr_pct"] >= 2)].sort_values("max_dd_pct", ascending=False).head(12)
    for _, r in low.iterrows():
        lines.append(
            f"  [{r['family']}] {r['name']}: Sharpe {r['sharpe']:.3f}  "
            f"CAGR {r['cagr_pct']:.1f}%  MaxDD {r['max_dd_pct']:.1f}%"
        )

    cmd = (
        f"cd {_REPO} && PYTHONUNBUFFERED=1 .venv/bin/python "
        f"RenTech/strategy_stack/discover_french_decile_strategies.py "
        f"--weighting {weighting} --min-sharpe {min_sharpe} --out-prefix {out_prefix}"
    )
    if start:
        cmd += f" --start {start}"
    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)
    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("--min-sharpe", type=float, default=0.25)
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--out-prefix", type=Path, default=LOGS / "french_decile_discovery_vw")
    ap.add_argument("-q", "--quiet", action="store_true")
    args = ap.parse_args()
    discover(
        start=args.start,
        end=args.end,
        weighting=args.weighting,
        min_sharpe=args.min_sharpe,
        capital=args.capital,
        out_prefix=args.out_prefix,
        verbose=not args.quiet,
    )


if __name__ == "__main__":
    main()
