"""
Tier A: **index / ETF total-return proxies** plus **documented synthetics** when Yahoo has no index.

Synthetics are **not** CBOE marks — they are rule-based series built from observables (SPY, SHV, VIX,
BXM, PUTW, VIXY, SVXY, UVXY, ^VVIX) so you can **compare shapes** (CAGR, vol, max DD, regime slices)
before wiring full option replication in Theta.

Requires: ``pip install yfinance pandas numpy`` (already typical for this repo).

**Chat-topic → row id (this module does not implement full option legs):**

- Long strangle / tail / VXTH-ish → ``syn_straddle_complacency``, ``syn_vxth_style``, ``vixy``
- BXM / covered call / PMCC *theme* → ``bxm``, ``pbp``
- Put-write / PUTY → ``putw_etf`` (ETF proxy; CBOE PUT index N/A on Yahoo)
- Iron condor → ``syn_iron_condor_blend``
- Jade lizard → ``syn_jade_lizard_blend``
- 0DTE → ``syn_zero_dte_vrp``
- Short vol 12–20 → ``syn_short_vol_normal``, ``svxy``
- R3-style wider PCS / CSP *theme* → ``putw_etf`` / ``bxm`` (imperfect), ``syn_vrp_carry``
- Panic / long vol → ``uvxy``, ``syn_long_uvxy_panic``
- Calendar / vol-of-vol → ``syn_calendar_vvix``
- VRP / term / skew **filters** → not separate rows; partial intent in ``syn_vrp_carry`` / VVIX tilt
"""

from __future__ import annotations

import math
from typing import Any

import numpy as np
import pandas as pd


def yf_adj_close(ticker: str, start: str, end: str) -> pd.Series:
    """Download adjusted close (``auto_adjust=True``) as a single float Series indexed by date."""
    try:
        import yfinance as yf
    except ImportError as e:
        raise ImportError("yfinance required: pip install yfinance") from e

    raw = yf.download(
        ticker,
        start=start,
        end=end,
        progress=False,
        auto_adjust=True,
        threads=False,
    )
    if raw is None or raw.empty:
        return pd.Series(dtype=float, name=ticker)

    if isinstance(raw.columns, pd.MultiIndex):
        if "Close" in raw.columns.get_level_values(0):
            s = raw["Close"].iloc[:, 0]
        elif "Adj Close" in raw.columns.get_level_values(0):
            s = raw["Adj Close"].iloc[:, 0]
        else:
            s = raw.iloc[:, 0]
    else:
        s = raw["Close"] if "Close" in raw.columns else raw.iloc[:, 0]
    s = pd.to_numeric(s.squeeze(), errors="coerce").astype(float)
    s.index = pd.to_datetime(s.index).normalize()
    s.name = ticker
    return s.sort_index().dropna()


def prices_to_returns(close: pd.Series) -> pd.Series:
    r = close.pct_change()
    r = pd.to_numeric(r, errors="coerce")
    return r.replace([np.inf, -np.inf], np.nan).dropna()


def wealth_index(daily_ret: pd.Series, start: float = 1.0) -> pd.Series:
    return (1.0 + daily_ret.fillna(0.0)).cumprod() * start


def max_drawdown(wealth: pd.Series) -> float:
    if wealth.empty:
        return float("nan")
    w = wealth.astype(float)
    peak = w.cummax()
    dd = (w - peak) / peak.replace(0.0, np.nan)
    return float(-dd.min()) if math.isfinite(-dd.min()) else float("nan")


def cagr_from_wealth(wealth: pd.Series) -> float:
    if wealth.size < 2:
        return float("nan")
    w0, w1 = float(wealth.iloc[0]), float(wealth.iloc[-1])
    if w0 <= 0 or w1 <= 0:
        return float("nan")
    days = (wealth.index[-1] - wealth.index[0]).days
    if days < 30:
        return float("nan")
    years = days / 365.25
    return float((w1 / w0) ** (1.0 / years) - 1.0)


def ann_vol(daily_ret: pd.Series) -> float:
    r = daily_ret.dropna().astype(float)
    if r.size < 20:
        return float("nan")
    return float(r.std(ddof=0) * math.sqrt(252.0))


def sharpe(daily_ret: pd.Series, rf_daily: pd.Series | None = None) -> float:
    r = daily_ret.dropna().astype(float)
    if r.size < 20:
        return float("nan")
    if rf_daily is not None:
        rf = rf_daily.reindex(r.index).dropna()
        x = (r - rf).reindex(rf.index).dropna()
    else:
        x = r
    if x.size < 20 or x.std(ddof=0) < 1e-12:
        return float("nan")
    return float((x.mean() / x.std(ddof=0)) * math.sqrt(252.0))


def regime_conditional_means(daily_ret: pd.Series, vix: pd.Series) -> dict[str, float]:
    """Mean **daily** return in each VIX bucket (labels match VRP engine: R1 <12, R2 12–20, …)."""
    df = pd.DataFrame({"r": daily_ret, "vx": vix}).dropna()
    out: dict[str, float] = {}
    m0 = df.loc[df["vx"] < 12.0, "r"]
    out["lt_12"] = float(m0.mean()) if len(m0) else float("nan")
    m1 = df.loc[(df["vx"] >= 12.0) & (df["vx"] <= 20.0), "r"]
    out["12_20"] = float(m1.mean()) if len(m1) else float("nan")
    m2 = df.loc[(df["vx"] > 20.0) & (df["vx"] <= 30.0), "r"]
    out["20_30"] = float(m2.mean()) if len(m2) else float("nan")
    m3 = df.loc[df["vx"] > 30.0, "r"]
    out["gt_30"] = float(m3.mean()) if len(m3) else float("nan")
    return out


def inner_join_returns(
    series: dict[str, pd.Series],
    *,
    min_overlap: int = 200,
) -> tuple[pd.DataFrame, list[str]]:
    """Align all series on intersection of dates; return log of dropped keys if too short."""
    if not series:
        return pd.DataFrame(), []
    common: pd.Index | None = None
    for _, s in series.items():
        idx = s.dropna().index
        common = idx if common is None else common.intersection(idx)
    if common is None or len(common) < min_overlap:
        return pd.DataFrame(), list(series.keys())
    common = common.sort_values()
    out = pd.DataFrame({k: v.reindex(common) for k, v in series.items()}).dropna(how="any")
    dropped = [k for k in series if k not in out.columns]
    return out, dropped


# --- Synthetics (documented) -------------------------------------------------


def synthetic_vrp_variance_carry(
    spy_ret: pd.Series,
    shv_ret: pd.Series,
    vix: pd.Series,
    *,
    k: float = 0.32,
    vrp_clip: float = 0.08,
) -> pd.Series:
    """
    Short-variance **carry toy**: earn when implied variance (VIX²) exceeds trailing 20d realized
    variance of SPY; lose otherwise. Scaled by ``k`` and clipped — **not** an options PnL path.
    """
    rv = spy_ret.rolling(20, min_periods=15).std() * math.sqrt(252.0)
    iv = (vix / 100.0).astype(float)
    vrp = (iv**2 - rv**2).clip(-vrp_clip, vrp_clip)
    carry = k * vrp / 252.0
    return (1.0 + shv_ret) * (1.0 + carry) - 1.0


def synthetic_straddle_complacency(
    spy_ret: pd.Series,
    shv_ret: pd.Series,
    vix: pd.Series,
    *,
    vix_enter: float = 12.0,
    gamma_scale: float = 6.0,
    theta_per_day: float = 0.35,
) -> pd.Series:
    """
    **Toy** long-straddle PnL when VIX < ``vix_enter``: small daily theta drain + payout when
    |SPY| move exceeds a VIX-implied daily sigma. Else hold SHV. For **shape** vs BXM only.
    """
    sig_d = (vix / 100.0) / math.sqrt(252.0)
    move = spy_ret.abs()
    payout = gamma_scale * (move - sig_d).clip(lower=0.0)
    theta = -theta_per_day * (vix / 100.0) / math.sqrt(252.0)
    long_leg = theta / 252.0 + payout
    mask = (vix < vix_enter).astype(float)
    blend = (mask * long_leg + (1.0 - mask) * 0.0).clip(-0.12, 0.12)
    return (1.0 + shv_ret) * (1.0 + blend) - 1.0


def synthetic_vxth_style_vixy(
    vixy_ret: pd.Series,
    shv_ret: pd.Series,
    vix: pd.Series,
    *,
    vix_ref: float = 15.0,
    width: float = 10.0,
) -> pd.Series:
    """Weight VIXY higher when VIX is below ``vix_ref`` (VXTH-style *intent*); daily rebalance."""
    w = ((vix_ref - vix) / width).clip(0.0, 1.0).astype(float)
    return w * vixy_ret + (1.0 - w) * shv_ret


def synthetic_short_vol_slice(
    svxy_ret: pd.Series,
    shv_ret: pd.Series,
    vix: pd.Series,
    *,
    lo: float = 12.0,
    hi: float = 20.0,
    weight: float = 0.12,
) -> pd.Series:
    """Small SVXY sleeve in the 12–20 band; **SVXY had a termination event (Feb 2018)** — use for
    historical comparison only."""
    m = ((vix >= lo) & (vix <= hi)).astype(float) * weight
    return m * svxy_ret + (1.0 - m) * shv_ret


def synthetic_long_uvxy_panic_slice(
    uvxy_ret: pd.Series,
    shv_ret: pd.Series,
    vix: pd.Series,
    *,
    thr: float = 30.0,
    weight: float = 0.10,
) -> pd.Series:
    """Small long-UVXY slice when VIX > ``thr`` (panic convexity toy); else SHV."""
    m = (vix > thr).astype(float) * weight
    return m * uvxy_ret + (1.0 - m) * shv_ret


def synthetic_iron_condor_blend(bxm_ret: pd.Series, shv_ret: pd.Series, w_bxm: float = 0.42) -> pd.Series:
    """No IC index on Yahoo: blend BXM (short premium index) with SHV as a **rough** IC-like risk."""
    return w_bxm * bxm_ret + (1.0 - w_bxm) * shv_ret


def synthetic_jade_lizard_blend(putw_ret: pd.Series, bxm_ret: pd.Series) -> pd.Series:
    """PUTW (put income) + BXM (call overwrite) blend — **intent** jade lizard (mixed calls/puts)."""
    return 0.5 * putw_ret + 0.5 * bxm_ret


def synthetic_zero_dte_vrp_aggressive(
    spy_ret: pd.Series,
    shv_ret: pd.Series,
    vix: pd.Series,
    *,
    k: float = 1.15,
    window: int = 5,
) -> pd.Series:
    """Aggressive short-horizon VRP toy: ``k`` × max(0, implied daily var − realized 5d var)."""
    rv5 = spy_ret.rolling(window, min_periods=3).std() * math.sqrt(252.0)
    iv = (vix / 100.0).astype(float)
    gap = (iv**2 - rv5**2).clip(-0.25, 0.25) / 252.0
    carry = k * gap
    return (1.0 + shv_ret) * (1.0 + carry) - 1.0


def synthetic_calendar_vvix_tilt(
    vixy_ret: pd.Series,
    shv_ret: pd.Series,
    vvix: pd.Series,
    *,
    zwin: int = 20,
    tilt: float = 0.18,
) -> pd.Series:
    """Tilt toward VIXY when ^VVIX is above its rolling median (vol-of-vol / calendar *intent*)."""
    med = vvix.rolling(zwin, min_periods=max(5, zwin // 4)).median()
    z = ((vvix - med) / (med + 1e-6)).clip(-3.0, 3.0)
    w = (tilt * (z > 0.0).astype(float) * z.clip(0.0, 1.0)).clip(0.0, 0.45)
    return w * vixy_ret + (1.0 - w) * shv_ret


def build_return_panel(
    start: str,
    end: str,
    *,
    exclude_optional: bool = False,
) -> tuple[pd.DataFrame, pd.Series, dict[str, str]]:
    """
    Fetch raw closes, build **daily returns** for Yahoo symbols + synthetics.

    Parameters
    ----------
    exclude_optional
        If True, omit PUTW / PBP so the inner-join starts at **2008** (BXM era) instead of ETF
        inception (~2011–2016). Synthetics that need PUTW are omitted in that mode.

    Returns
    -------
    returns_panel, vix_level (aligned to panel index), warnings
    """
    warnings: dict[str, str] = {}
    tickers: dict[str, str] = {
        "sp500_tr": "^SP500TR",
        "spy": "SPY",
        "shv": "SHV",
        "bxm": "^BXM",
        "vixy": "VIXY",
        "svxy": "SVXY",
        "uvxy": "UVXY",
        "vix": "^VIX",
        "vvix": "^VVIX",
    }
    if not exclude_optional:
        tickers["putw"] = "PUTW"
        tickers["pbp"] = "PBP"
        # SPYC (2020 inception) is omitted here — it collapses the whole panel to ~800 days.
        # Use PBP / BXM for buy-write PMCC *theme* on the long sample.

    closes: dict[str, pd.Series] = {}
    for key, tkr in tickers.items():
        s = yf_adj_close(tkr, start, end)
        if s.empty:
            warnings[key] = f"empty series for {tkr}"
        closes[key] = s

    joined, dropped = inner_join_returns(closes, min_overlap=120)
    for k in dropped:
        warnings[f"join_{k}"] = "dropped (insufficient overlap)"

    if joined.empty:
        return pd.DataFrame(), pd.Series(dtype=float), warnings

    rets = joined.pct_change().replace([np.inf, -np.inf], np.nan).dropna(how="any")
    vix_lvl = joined["vix"].shift(1).reindex(rets.index).ffill()
    vvix_lvl = joined["vvix"].shift(1).reindex(rets.index).ffill()

    spy_r = rets["spy"]
    shv_r = rets["shv"]
    bxm_r = rets["bxm"]
    vixy_r = rets["vixy"]
    svxy_r = rets["svxy"]
    uvxy_r = rets["uvxy"]

    cols: dict[str, pd.Series] = {
        "sp500_tr": rets["sp500_tr"],
        "spy_price": rets["spy"],
        "shv_cash": rets["shv"],
        "bxm": rets["bxm"],
        "vixy": rets["vixy"],
        "svxy": rets["svxy"],
        "uvxy": rets["uvxy"],
        "syn_vrp_carry": synthetic_vrp_variance_carry(spy_r, shv_r, vix_lvl),
        "syn_straddle_complacency": synthetic_straddle_complacency(spy_r, shv_r, vix_lvl),
        "syn_vxth_style": synthetic_vxth_style_vixy(vixy_r, shv_r, vix_lvl),
        "syn_short_vol_normal": synthetic_short_vol_slice(svxy_r, shv_r, vix_lvl),
        "syn_long_uvxy_panic": synthetic_long_uvxy_panic_slice(uvxy_r, shv_r, vix_lvl),
        "syn_iron_condor_blend": synthetic_iron_condor_blend(bxm_r, shv_r),
        "syn_zero_dte_vrp": synthetic_zero_dte_vrp_aggressive(spy_r, shv_r, vix_lvl),
        "syn_calendar_vvix": synthetic_calendar_vvix_tilt(vixy_r, shv_r, vvix_lvl),
    }
    if not exclude_optional and "putw" in rets.columns:
        putw_r = rets["putw"]
        cols["putw_etf"] = putw_r
        cols["syn_jade_lizard_blend"] = synthetic_jade_lizard_blend(putw_r, bxm_r)
    if not exclude_optional and "pbp" in rets.columns:
        cols["pbp"] = rets["pbp"]

    out = pd.DataFrame(cols)
    out = out.replace([np.inf, -np.inf], np.nan).dropna(how="any")
    vx = vix_lvl.reindex(out.index).ffill()
    return out, vx, warnings


# Columns allowed on :class:`~RenTech.strategy_stack.vrp_backtester.VRPBacktester` macro overlay.
MACRO_OVERLAY_ALLOWED_IDS: frozenset[str] = frozenset(
    {
        "syn_straddle_complacency",
        "syn_zero_dte_vrp",
        "syn_jade_lizard_blend",
        "putw_etf",
        "bxm",
        "pbp",
        "syn_iron_condor_blend",
        "syn_short_vol_normal",
        "syn_vrp_carry",
    }
)

# Default mapping: Tier A / ETF sleeves by VIX band (see module docstring in vrp_backtester).
DEFAULT_MACRO_OVERLAY_BY_BAND: dict[str, tuple[str, ...]] = {
    # R1 complacency: long-convexity / straddle *intent* (synthetic); not short premium.
    "r1": ("syn_straddle_complacency",),
    # R2 normal: short-premium + VRP harvest + buy-write / put-write proxies; SVXY slice is 12–20 only by construction.
    "r2": (
        "syn_zero_dte_vrp",
        "syn_jade_lizard_blend",
        "putw_etf",
        "bxm",
        "pbp",
        "syn_iron_condor_blend",
        "syn_short_vol_normal",
        "syn_vrp_carry",
    ),
    # R3 elevated: keep short-premium theme; drop dedicated 12–20 SVXY sleeve (wrong band).
    "r3": (
        "putw_etf",
        "bxm",
        "pbp",
        "syn_jade_lizard_blend",
        "syn_iron_condor_blend",
        "syn_vrp_carry",
        "syn_zero_dte_vrp",
    ),
    # R4 panic: only mild carry / defensive blend — no jade, no 0DTE toy, no SVXY short vol.
    "r4": ("syn_iron_condor_blend", "syn_vrp_carry"),
}


def load_macro_returns_aligned(
    spy_df: pd.DataFrame,
    *,
    exclude_optional_etf: bool = False,
) -> pd.DataFrame:
    """
    Build the Tier A **daily return** columns on ``spy_df``'s calendar (forward-filled from first
    valid Yahoo overlap). Used by :class:`~RenTech.strategy_stack.vrp_backtester.VRPBacktester`
    macro overlay.
    """
    if spy_df is None or spy_df.empty:
        return pd.DataFrame()
    idx = pd.DatetimeIndex(pd.to_datetime(spy_df.index).normalize()).sort_values()
    start = idx.min().strftime("%Y-%m-%d")
    end = (idx.max() + pd.Timedelta(days=2)).strftime("%Y-%m-%d")
    panel, _, _ = build_return_panel(start, end, exclude_optional=exclude_optional_etf)
    if panel.empty:
        return pd.DataFrame()
    out = panel.reindex(idx)
    return out.ffill()


def summarize_strategy(
    daily_ret: pd.Series,
    vix_level: pd.Series,
    rf_daily: pd.Series | None,
) -> dict[str, Any]:
    w = wealth_index(daily_ret)
    reg = regime_conditional_means(daily_ret, vix_level.reindex(daily_ret.index).ffill())
    return {
        "n_days": int(daily_ret.shape[0]),
        "total_return": float(w.iloc[-1] / w.iloc[0] - 1.0) if len(w) > 1 else float("nan"),
        "cagr": cagr_from_wealth(w),
        "ann_vol": ann_vol(daily_ret),
        "max_drawdown": max_drawdown(w),
        "sharpe": sharpe(daily_ret, rf_daily),
        "mean_daily_lt12": reg["lt_12"],
        "mean_daily_12_20": reg["12_20"],
        "mean_daily_20_30": reg["20_30"],
        "mean_daily_gt30": reg["gt_30"],
    }
