#!/usr/bin/env python3
"""
Walk-forward Markov-chain backtest on the **All-Weather** macro basket.

Universe (Bridgewater-style baseline from :mod:`portfolio_risk_manager`):
  SPY 30% · TLT 40% · IEF 15% · GLD 7.5% · DBC 7.5%

Each sleeve gets a rolling-percentile Markov state machine. Per day:
  1. Re-estimate transition matrix from trailing history (walk-forward)
  2. Monte Carlo → calibrated probability of ending "bullish"
  3. Compare vs current percentile (market price analog)
  4. Quarter-Kelly exposure scales baseline weight (0 = cash)

Benchmarks in output:
  * **markov_aw** — dynamic Markov-gated weights
  * **static_aw** — always fully invested at BASE_WEIGHTS
  * **tactical_aw** — momentum + SMA200 gates (:class:`TacticalAllWeatherManager`)

Example::

    cd /Users/robzingale/trading_bot
    PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_markov_all_weather_backtest.py \\
        --start 2016-01-04 --yahoo-period max
"""

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.data_loader import DataLoader
from RenTech.strategy_stack.main import _compute_daily_backtest_features
from RenTech.strategy_stack.markov_chain_trading import (
    MarkovChainTradingModel,
    walk_forward_exposure_series,
)
from RenTech.strategy_stack.portfolio_risk_manager import BASE_WEIGHTS, TacticalAllWeatherManager

LOGS = _REPO / "RenTech" / "data" / "logs"
DEFAULT_OUT_PREFIX = LOGS / "markov_all_weather"
MACRO_TICKERS = tuple(BASE_WEIGHTS.keys())


def _load_macro_dict(yahoo_period: str) -> dict[str, pd.DataFrame]:
    loader = DataLoader()
    out: dict[str, pd.DataFrame] = {}
    for t in MACRO_TICKERS:
        daily = loader.fetch_daily(t, period=yahoo_period)
        if daily.empty:
            raise RuntimeError(f"No daily data for {t}")
        out[t] = _compute_daily_backtest_features(daily)
    return out


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


def _load_vix_series(yahoo_period: str) -> pd.Series:
    loader = DataLoader()
    vix_df = loader.fetch_daily("^VIX", period=yahoo_period)
    if vix_df.empty:
        raise RuntimeError("No VIX data from ^VIX")
    s = vix_df["close"].astype(np.float64)
    s.index = pd.to_datetime(s.index).tz_localize(None)
    return s


def build_markov_portfolio(
    macro_dict: dict[str, pd.DataFrame],
    *,
    model: MarkovChainTradingModel,
    percentile_window: int,
    lookback: int,
    cash_annual_yield: float,
    matrix_refresh: int = 5,
    vix: pd.Series | None = None,
    vix_lo: float = 15.0,
    vix_hi: float = 25.0,
) -> pd.DataFrame:
    """Combine per-ticker Markov exposures with BASE_WEIGHTS."""
    sleeves = list(BASE_WEIGHTS.keys())
    indices = [
        pd.to_datetime(macro_dict[t].index).tz_localize(None).sort_values()
        for t in sleeves
    ]
    master = indices[0]
    for idx in indices[1:]:
        master = master.union(idx)
    master = master.sort_values()

    ret_df = pd.DataFrame(index=master, dtype=np.float64)
    exp_df = pd.DataFrame(index=master, dtype=np.float64)

    signal_frames: dict[str, pd.DataFrame] = {}
    for t in sleeves:
        df = macro_dict[t].copy()
        df.index = pd.to_datetime(df.index).tz_localize(None)
        close = df["close"].reindex(master).ffill()
        ret = df["ret"].reindex(master).fillna(0.0)
        ret_df[t] = ret

        sig = walk_forward_exposure_series(
            close,
            percentile_window=percentile_window,
            lookback=lookback,
            model=model,
            matrix_refresh=matrix_refresh,
            vix=vix,
            vix_lo=vix_lo,
            vix_hi=vix_hi,
        )
        signal_frames[t] = sig
        exp = sig["exposure"].reindex(master).fillna(0.0)
        exp_df[t] = exp

    base_w = pd.Series(BASE_WEIGHTS)
    target_weight = exp_df.mul(base_w, axis=1)
    total_invested = target_weight.sum(axis=1)
    cash_weight = 1.0 - total_invested
    daily_rf = cash_annual_yield / 252.0

    portfolio_ret = (ret_df * target_weight).sum(axis=1) + cash_weight * daily_rf
    static_ret = (ret_df * base_w).sum(axis=1)

    out = pd.DataFrame(
        {
            "portfolio_bar_ret": portfolio_ret,
            "static_all_weather_ret": static_ret,
            "cash_weight": cash_weight,
            "total_invested_weight": total_invested,
        },
        index=master,
    )
    for t in sleeves:
        out[f"weight_{t}"] = target_weight[t]
        out[f"exposure_{t}"] = exp_df[t]
        out[f"edge_{t}"] = signal_frames[t]["edge"].reindex(master)
    return out


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("--cash-yield", type=float, default=0.04)
    ap.add_argument("--n-states", type=int, default=10)
    ap.add_argument("--horizon", type=int, default=21, help="MC steps (trading days)")
    ap.add_argument("--n-sims", type=int, default=10_000)
    ap.add_argument("--percentile-window", type=int, default=252)
    ap.add_argument("--lookback", type=int, default=252, help="Transition-matrix history")
    ap.add_argument("--edge-threshold", type=float, default=0.03)
    ap.add_argument(
        "--calibration",
        choices=("none", "equity_shrink", "polymarket"),
        default="equity_shrink",
    )
    ap.add_argument("--kelly-mult", type=float, default=0.25)
    ap.add_argument(
        "--matrix-refresh",
        type=int,
        default=5,
        help="Re-estimate transition matrix every N days (default 5)",
    )
    ap.add_argument(
        "--vix-regime",
        action="store_true",
        help="Use separate transition matrices when VIX < lo vs VIX > hi",
    )
    ap.add_argument("--vix-lo", type=float, default=15.0, help="Low-vol regime threshold")
    ap.add_argument("--vix-hi", type=float, default=25.0, help="High-vol regime threshold")
    ap.add_argument(
        "--compare-pooled",
        action="store_true",
        help="Also run pooled (non-VIX) Markov book for side-by-side metrics",
    )
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT_PREFIX)
    args = ap.parse_args()

    model = MarkovChainTradingModel(
        n_states=args.n_states,
        horizon=args.horizon,
        n_sims=args.n_sims,
        kelly_mult=args.kelly_mult,
        edge_threshold=args.edge_threshold,
        calibration_mode=args.calibration,
    )

    print("Loading All-Weather macro data …", flush=True)
    macro_dict = _load_macro_dict(args.yahoo_period)

    vix_series: pd.Series | None = None
    if args.vix_regime:
        print("Loading VIX …", flush=True)
        vix_series = _load_vix_series(args.yahoo_period)

    print("Building Markov walk-forward portfolio …", flush=True)
    markov_port = build_markov_portfolio(
        macro_dict,
        model=model,
        percentile_window=args.percentile_window,
        lookback=args.lookback,
        cash_annual_yield=float(args.cash_yield),
        matrix_refresh=args.matrix_refresh,
        vix=vix_series if args.vix_regime else None,
        vix_lo=args.vix_lo,
        vix_hi=args.vix_hi,
    )

    markov_pooled_port: pd.DataFrame | None = None
    if args.vix_regime and args.compare_pooled:
        print("Building pooled Markov baseline for comparison …", flush=True)
        markov_pooled_port = build_markov_portfolio(
            macro_dict,
            model=model,
            percentile_window=args.percentile_window,
            lookback=args.lookback,
            cash_annual_yield=float(args.cash_yield),
            matrix_refresh=args.matrix_refresh,
            vix=None,
        )

    print("Building Tactical All-Weather benchmark …", flush=True)
    pm = TacticalAllWeatherManager()
    tactical_port = pm.build_portfolio(macro_dict, cash_annual_yield=float(args.cash_yield))
    tactical_port.index = pd.to_datetime(tactical_port.index).tz_localize(None)

    markov_port.index = pd.to_datetime(markov_port.index).tz_localize(None)
    combined = markov_port.join(
        tactical_port[["portfolio_bar_ret"]].rename(
            columns={"portfolio_bar_ret": "tactical_aw_ret"}
        ),
        how="left",
    )
    if markov_pooled_port is not None and args.vix_regime:
        markov_pooled_port.index = pd.to_datetime(markov_pooled_port.index).tz_localize(None)
        combined = combined.join(
            markov_pooled_port[["portfolio_bar_ret"]].rename(
                columns={"portfolio_bar_ret": "markov_pooled_ret"}
            ),
            how="left",
        )

    mask = combined.index >= pd.Timestamp(args.start)
    if args.end.strip():
        mask &= combined.index <= pd.Timestamp(args.end)
    win = combined.loc[mask].fillna(0.0)

    cap = float(args.capital)
    r_markov = win["portfolio_bar_ret"].astype(np.float64)
    r_static = win["static_all_weather_ret"].astype(np.float64)
    r_tactical = win["tactical_aw_ret"].astype(np.float64)
    r_pooled = (
        win["markov_pooled_ret"].astype(np.float64)
        if "markov_pooled_ret" in win.columns
        else None
    )

    m_markov = _metrics(r_markov, cap)
    m_static = _metrics(r_static, cap)
    m_tactical = _metrics(r_tactical, cap)
    m_pooled = _metrics(r_pooled, cap) if r_pooled is not None else None

    spy_ret = macro_dict["SPY"]["ret"].astype(float)
    spy_ret.index = pd.to_datetime(spy_ret.index).tz_localize(None)
    corr_df = pd.DataFrame({"markov": r_markov, "spy": spy_ret}).dropna()
    rho_spy = float(corr_df.corr().iloc[0, 1]) if len(corr_df) > 5 else float("nan")

    strategy_label = "markov_vix_regime" if args.vix_regime else "markov_chain_all_weather"
    if args.vix_regime:
        prefix = Path(str(args.out_prefix) + "_vix_regime").expanduser().resolve()
    else:
        prefix = args.out_prefix.expanduser().resolve()
    prefix.parent.mkdir(parents=True, exist_ok=True)
    daily_path = Path(f"{prefix}_daily.csv")
    meta_path = Path(f"{prefix}_meta.json")
    metrics_path = Path(f"{prefix}_metrics.txt")
    signals_path = Path(f"{prefix}_signals.csv")

    eq_markov = cap * (1.0 + r_markov).cumprod()
    out_daily = pd.DataFrame(
        {
            "date": r_markov.index.strftime("%Y-%m-%d"),
            "daily_ret_markov": r_markov.values,
            "daily_ret_static_aw": r_static.values,
            "daily_ret_tactical_aw": r_tactical.values,
            "equity_markov_usd": eq_markov.values,
            "cash_weight": win["cash_weight"].values,
            "total_invested_weight": win["total_invested_weight"].values,
        }
    )
    if r_pooled is not None:
        out_daily["daily_ret_markov_pooled"] = r_pooled.values
    out_daily.to_csv(daily_path, index=False)

    sig_cols = [c for c in win.columns if c.startswith(("exposure_", "edge_", "weight_"))]
    sig_out = win[sig_cols].copy()
    sig_out.insert(0, "date", sig_out.index.strftime("%Y-%m-%d"))
    sig_out.to_csv(signals_path, index=False)

    start_d = str(r_markov.index.min().date()) if len(r_markov) else None
    end_d = str(r_markov.index.max().date()) if len(r_markov) else None

    cmd = (
        f"cd {_REPO} && PYTHONUNBUFFERED=1 .venv/bin/python "
        f"RenTech/strategy_stack/run_markov_all_weather_backtest.py "
        f"--start {args.start} --yahoo-period {args.yahoo_period} "
        f"--calibration {args.calibration}"
    )
    if args.end.strip():
        cmd += f" --end {args.end}"

    if args.vix_regime:
        cmd += f" --vix-regime --vix-lo {args.vix_lo} --vix-hi {args.vix_hi}"
    if args.compare_pooled:
        cmd += " --compare-pooled"

    meta = {
        "strategy": strategy_label,
        "baseline_weights": BASE_WEIGHTS,
        "model": {
            "n_states": args.n_states,
            "horizon": args.horizon,
            "n_sims": args.n_sims,
            "percentile_window": args.percentile_window,
            "lookback": args.lookback,
            "edge_threshold": args.edge_threshold,
            "calibration": args.calibration,
            "kelly_mult": args.kelly_mult,
            "matrix_refresh": args.matrix_refresh,
            "vix_regime": bool(args.vix_regime),
            "vix_lo": args.vix_lo,
            "vix_hi": args.vix_hi,
        },
        "start": start_d,
        "end": end_d,
        "n_days": int(len(r_markov)),
        "capital": cap,
        "cash_annual_yield": float(args.cash_yield),
        "markov_aw": m_markov,
        "markov_pooled": m_pooled,
        "static_all_weather": m_static,
        "tactical_all_weather": m_tactical,
        "corr_markov_vs_spy": rho_spy,
        "artifacts": {
            "daily_csv": str(daily_path),
            "signals_csv": str(signals_path),
            "meta_json": str(meta_path),
        },
    }
    meta_path.write_text(json.dumps(meta, indent=2) + "\n")

    metrics_lines = [
        f"# Markov Chain All-Weather backtest ({start_d} -> {end_d})\n",
        f"Command:\n{cmd}\n",
        f"Universe: {BASE_WEIGHTS}\n",
    ]
    if args.vix_regime:
        metrics_lines.append(
            f"VIX regimes: low < {args.vix_lo}, high > {args.vix_hi}, mid = pooled mid matrix\n"
        )

    def _block(title: str, m: dict[str, float]) -> str:
        return (
            f"--- {title} ---\n"
            f"  Total return: {m['total_return_pct']:.2f}%\n"
            f"  CAGR: {m['cagr_pct']:.2f}%\n"
            f"  Sharpe: {m['sharpe']:.3f}\n"
            f"  Max DD: {m['max_drawdown_pct']:.2f}%\n\n"
        )

    markov_title = "Markov AW (VIX regime)" if args.vix_regime else "Markov AW (dynamic)"
    metrics_lines.append(_block(markov_title, m_markov))
    if m_pooled is not None and args.vix_regime:
        metrics_lines.append(_block("Markov AW (pooled baseline)", m_pooled))
    metrics_lines.append(_block("Static All-Weather (benchmark)", m_static))
    metrics_lines.append(_block("Tactical All-Weather (benchmark)", m_tactical))
    metrics_lines.append(f"  Corr(Markov, SPY): {rho_spy:.3f}\n\n")
    metrics_lines.append(f"Artifacts:\n  {daily_path}\n  {signals_path}\n  {meta_path}\n")
    metrics_path.write_text("".join(metrics_lines))

    print(f"\n# Markov Chain All-Weather ({start_d} -> {end_d}, {len(r_markov)} days)\n")
    print(f"Command:\n{cmd}\n")
    if args.vix_regime:
        print(f"VIX regimes: low < {args.vix_lo}, high > {args.vix_hi}\n")
    print(f"--- {markov_title} ---")
    print(f"  Total return: {m_markov['total_return_pct']:.1f}%")
    print(f"  CAGR: {m_markov['cagr_pct']:.2f}%")
    print(f"  Sharpe: {m_markov['sharpe']:.3f}")
    print(f"  Max DD: {m_markov['max_drawdown_pct']:.2f}%")
    if m_pooled is not None and args.vix_regime:
        print("\n--- Markov AW (pooled baseline) ---")
        print(f"  Total return: {m_pooled['total_return_pct']:.1f}%")
        print(f"  CAGR: {m_pooled['cagr_pct']:.2f}%")
        print(f"  Sharpe: {m_pooled['sharpe']:.3f}")
        print(f"  Max DD: {m_pooled['max_drawdown_pct']:.2f}%")
        delta_ret = m_markov["total_return_pct"] - m_pooled["total_return_pct"]
        delta_sh = m_markov["sharpe"] - m_pooled["sharpe"]
        print(f"\n  VIX regime vs pooled: return {delta_ret:+.1f}pp, Sharpe {delta_sh:+.3f}")
    print("\n--- Static All-Weather ---")
    print(f"  Total return: {m_static['total_return_pct']:.1f}%")
    print(f"  Sharpe: {m_static['sharpe']:.3f}")
    print(f"  Max DD: {m_static['max_drawdown_pct']:.2f}%")
    print("\n--- Tactical All-Weather ---")
    print(f"  Total return: {m_tactical['total_return_pct']:.1f}%")
    print(f"  Sharpe: {m_tactical['sharpe']:.3f}")
    print(f"  Max DD: {m_tactical['max_drawdown_pct']:.2f}%")
    print(f"\nCorr(Markov, SPY): {rho_spy:.3f}")
    print(f"\nWrote {daily_path}")
    print(f"Wrote {signals_path}")
    print(f"Wrote {meta_path}")


if __name__ == "__main__":
    main()
