#!/usr/bin/env python3
"""
Grid-search Tactical All Weather gate/sizing variants (2016+).

Compares baseline vs partial weights, faster momentum, lower SMA, bond down-weight,
and SPY risk-on invested floor. Writes ranked CSV + markdown summary under
``RenTech/data/logs/``.

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 \\
        .venv/bin/python RenTech/strategy_stack/run_tactical_aw_variant_sweep.py \\
        --start 2016-01-04
"""
from __future__ import annotations

import argparse
import json
import sys
from dataclasses import asdict
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.portfolio_risk_manager import (
    BASE_WEIGHTS,
    TacticalAWConfig,
    TacticalAllWeatherManager,
)
from RenTech.strategy_stack.run_tactical_all_weather_standard import (
    MACRO_TICKERS,
    _load_macro_dict,
)

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


def _metrics(r: pd.Series, capital: float = 100_000.0) -> dict[str, float]:
    r = r.astype(np.float64).dropna()
    if len(r) < 2:
        return {}
    eq = capital * (1.0 + r).cumprod()
    years = len(r) / 252.0
    end = float(eq.iloc[-1])
    tot = end / capital - 1.0
    cagr = (end / capital) ** (1.0 / years) - 1.0 if years > 0 else float("nan")
    dd = float((eq / eq.cummax() - 1.0).min())
    sd = float(r.std(ddof=1))
    sharpe = float(r.mean() / sd * np.sqrt(252.0)) if sd > 1e-12 else float("nan")
    return {
        "total_return_pct": tot * 100.0,
        "cagr_pct": cagr * 100.0,
        "max_dd_pct": dd * 100.0,
        "sharpe": sharpe,
        "end_equity": end,
    }


def _period_return(r: pd.Series, start: str, end: str) -> float:
    sub = r.loc[start:end]
    if len(sub) < 2:
        return float("nan")
    return float((1.0 + sub).prod() - 1.0) * 100.0


def _avg_invested(port: pd.DataFrame, start: str, end: str) -> float:
    w = port["total_invested_weight"].loc[start:end]
    return float(w.mean()) * 100.0 if len(w) else float("nan")


def _run_variant(
    macro_dict: dict[str, pd.DataFrame],
    *,
    name: str,
    config: TacticalAWConfig,
    start: str,
    end: str,
    capital: float,
) -> dict:
    pm = TacticalAllWeatherManager(config=config)
    port = pm.build_portfolio(macro_dict)
    port.index = pd.to_datetime(port.index).tz_localize(None)
    mask = port.index >= pd.Timestamp(start)
    if end.strip():
        mask &= port.index <= pd.Timestamp(end)
    r = port.loc[mask, "portfolio_bar_ret"].astype(np.float64).fillna(0.0)
    rs = port.loc[mask, "static_all_weather_ret"] if "static_all_weather_ret" in port.columns else None
    if rs is None:
        base_w = pd.Series(BASE_WEIGHTS, dtype=np.float64)
        sleeves = list(BASE_WEIGHTS.keys())
        ret_df = pd.DataFrame(
            {t: macro_dict[t]["ret"].reindex(port.index).fillna(0.0) for t in sleeves}
        )
        rs = (ret_df * base_w).sum(axis=1).loc[mask]

    m_full = _metrics(r, capital)
    row = {
        "variant": name,
        **{f"cfg_{k}": v for k, v in asdict(config).items()},
        **m_full,
        "ret_2017_2021_pct": _period_return(r, "2017-01-01", "2021-12-31"),
        "ret_2022_2024_pct": _period_return(r, "2022-01-01", "2024-12-31"),
        "ret_2022_pct": _period_return(r, "2022-01-01", "2022-12-31"),
        "ret_2023_pct": _period_return(r, "2023-01-01", "2023-12-31"),
        "ret_2024_pct": _period_return(r, "2024-01-01", "2024-12-31"),
        "avg_invested_2022_2024_pct": _avg_invested(port, "2022-01-01", "2024-12-31"),
        "static_ret_2022_2024_pct": _period_return(rs.astype(float), "2022-01-01", "2024-12-31"),
    }
    return row


def _variant_grid() -> list[tuple[str, TacticalAWConfig]]:
    base = TacticalAWConfig()
    variants: list[tuple[str, TacticalAWConfig]] = [
        ("baseline_sma200_mom12_1", base),
        (
            "partial50_sma200_mom12_1",
            TacticalAWConfig(weight_mode="partial", partial_frac=0.5),
        ),
        (
            "partial50_sma100_mom12_1",
            TacticalAWConfig(sma_window=100, weight_mode="partial", partial_frac=0.5),
        ),
        (
            "mom6_1_thresh0",
            TacticalAWConfig(mom_lookback_days=126, mom_threshold=0.0),
        ),
        (
            "mom6_1_partial50",
            TacticalAWConfig(
                mom_lookback_days=126,
                weight_mode="partial",
                partial_frac=0.5,
            ),
        ),
        (
            "mom3_1_partial50",
            TacticalAWConfig(
                mom_lookback_days=63,
                weight_mode="partial",
                partial_frac=0.5,
            ),
        ),
        (
            "mom_thresh_m2",
            TacticalAWConfig(mom_threshold=-0.02),
        ),
        (
            "mom6_1_thresh_m2_partial50",
            TacticalAWConfig(
                mom_lookback_days=126,
                mom_threshold=-0.02,
                weight_mode="partial",
                partial_frac=0.5,
            ),
        ),
        (
            "bond70_sma200",
            TacticalAWConfig(bond_baseline_mult=0.7),
        ),
        (
            "bond70_partial50_mom6_1",
            TacticalAWConfig(
                bond_baseline_mult=0.7,
                mom_lookback_days=126,
                weight_mode="partial",
                partial_frac=0.5,
            ),
        ),
        (
            "risk_on_floor50",
            TacticalAWConfig(risk_on_min_invested=0.50),
        ),
        (
            "risk_on_floor50_mom6_1_partial50",
            TacticalAWConfig(
                risk_on_min_invested=0.50,
                mom_lookback_days=126,
                weight_mode="partial",
                partial_frac=0.5,
            ),
        ),
        (
            "combo_sma100_mom6_1_partial50_bond70_floor40",
            TacticalAWConfig(
                sma_window=100,
                mom_lookback_days=126,
                weight_mode="partial",
                partial_frac=0.5,
                bond_baseline_mult=0.7,
                risk_on_min_invested=0.40,
            ),
        ),
        (
            "combo_sma100_mom6_1_partial50_bond70",
            TacticalAWConfig(
                sma_window=100,
                mom_lookback_days=126,
                weight_mode="partial",
                partial_frac=0.5,
                bond_baseline_mult=0.7,
            ),
        ),
    ]
    return variants


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
    ap.add_argument("--start", default="2016-01-04")
    ap.add_argument("--end", default="")
    ap.add_argument("--yahoo-period", default="max")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument(
        "--out-prefix",
        type=Path,
        default=LOGS / "tactical_aw_variant_sweep",
    )
    args = ap.parse_args()

    macro_dict = _load_macro_dict(args.yahoo_period)
    rows = []
    for name, cfg in _variant_grid():
        print(f"  {name} …", flush=True)
        rows.append(
            _run_variant(
                macro_dict,
                name=name,
                config=cfg,
                start=str(args.start),
                end=str(args.end).strip(),
                capital=float(args.capital),
            )
        )

    df = pd.DataFrame(rows)
    df = df.sort_values(
        ["ret_2022_2024_pct", "sharpe"],
        ascending=[False, False],
    ).reset_index(drop=True)

    prefix = Path(args.out_prefix).expanduser().resolve()
    prefix.parent.mkdir(parents=True, exist_ok=True)
    csv_path = Path(f"{prefix}.csv")
    md_path = Path(f"{prefix}.md")
    json_path = Path(f"{prefix}_best.json")
    df.to_csv(csv_path, index=False)

    best = df.iloc[0]
    def _json_val(v: object) -> object:
        if isinstance(v, (np.floating, np.integer)):
            return float(v) if isinstance(v, np.floating) else int(v)
        return v

    cfg_out = {
        str(k).replace("cfg_", ""): _json_val(best[k])
        for k in df.columns
        if str(k).startswith("cfg_")
    }
    metrics_out = {
        k: _json_val(best[k])
        for k in (
            "total_return_pct",
            "cagr_pct",
            "max_dd_pct",
            "sharpe",
            "ret_2022_2024_pct",
            "ret_2017_2021_pct",
            "avg_invested_2022_2024_pct",
        )
    }
    json_path.write_text(
        json.dumps(
            {
                "best_variant": str(best["variant"]),
                "config": cfg_out,
                "metrics": metrics_out,
            },
            indent=2,
        )
        + "\n",
        encoding="utf-8",
    )

    baseline = df[df["variant"] == "baseline_sma200_mom12_1"].iloc[0]
    lines = [
        "# Tactical AW variant sweep",
        "",
        f"Window: {args.start} → {args.end or 'latest'} · capital ${args.capital:,.0f}",
        "",
        f"**Best by 2022–2024 return:** `{best['variant']}` — "
        f"2022–24 **{best['ret_2022_2024_pct']:+.1f}%** "
        f"(baseline {baseline['ret_2022_2024_pct']:+.1f}%), "
        f"full-window Sharpe **{best['sharpe']:.2f}**, max DD **{best['max_dd_pct']:.1f}%**, "
        f"avg invested 2022–24 **{best['avg_invested_2022_2024_pct']:.1f}%**.",
        "",
        "## Ranked by 2022–2024 return",
        "",
        df[
            [
                "variant",
                "ret_2022_2024_pct",
                "ret_2022_pct",
                "ret_2023_pct",
                "ret_2024_pct",
                "avg_invested_2022_2024_pct",
                "ret_2017_2021_pct",
                "sharpe",
                "max_dd_pct",
                "total_return_pct",
            ]
        ].to_string(index=False),
        "",
        f"CSV: `{csv_path}`",
    ]
    md_path.write_text("\n".join(lines) + "\n", encoding="utf-8")

    print(f"\nBest variant: {best['variant']}")
    print(
        f"  2022-24: {best['ret_2022_2024_pct']:+.1f}%  "
        f"(baseline {baseline['ret_2022_2024_pct']:+.1f}%)"
    )
    print(f"  Sharpe: {best['sharpe']:.2f}  MaxDD: {best['max_dd_pct']:.1f}%")
    print(f"\nWrote {csv_path}\nWrote {md_path}\nWrote {json_path}")


if __name__ == "__main__":
    main()
