"""
State-space features for SPX/SPY options regime routing (v2 spec).

Six dimensions: VRP, Hurst, vertical skew, vol term structure, spot-vol correlation,
spot-to-MA distance. Each is discretized into ordinal states for regime intersection.
"""

from __future__ import annotations

import math
from dataclasses import dataclass
from pathlib import Path
from typing import Literal

import numpy as np
import pandas as pd

VrpState = Literal["low", "normal", "high"]
HurstState = Literal["anti", "random", "persistent"]
SkewState = Literal["flat", "normal", "steep"]
TermState = Literal["contango", "backwardation"]
SpotVolState = Literal["negative", "positive"]
MaDistState = Literal["below", "neutral", "above"]

FeatureState = tuple[VrpState, HurstState, SkewState, TermState, SpotVolState, MaDistState]

# Thresholds from SPX_Regime_State_Space_v2.md
SKEW_FLAT_HI = 0.020
SKEW_NORMAL_HI = 0.035
MA_DIST_BAND = 0.01
HURST_ROLLING = 60
HURST_MAX_LAGS = 20
SPOT_VOL_WINDOW = 10
RV_WINDOW = 30
IV_TARGET_DTE = 30
SKEW_PUT_DELTA = -0.25
SKEW_CALL_DELTA = 0.25
MA_WINDOW = 50

_REPO = Path(__file__).resolve().parents[3]
CBOE_PANEL = _REPO / "RenTech" / "data" / "vix_futures_cboe.parquet"


def calculate_hurst_exponent(series: pd.Series, *, max_lags: int = HURST_MAX_LAGS) -> float:
    """Variance-of-lags Hurst estimate; <0.5 mean-reverting, >0.5 trending."""
    s = series.dropna().astype(np.float64)
    if len(s) < max_lags + 5:
        return float("nan")

    lags: list[int] = []
    tau: list[float] = []
    arr = s.to_numpy()
    for lag in range(2, int(max_lags) + 1):
        if lag >= len(arr):
            break
        diff = arr[lag:] - arr[:-lag]
        if diff.size < 5:
            continue
        v = np.var(diff, ddof=1)
        if v <= 0 or not np.isfinite(v):
            continue
        tau.append(float(np.sqrt(v)))
        lags.append(lag)

    if len(tau) < 3:
        return float("nan")

    slope, _ = np.polyfit(np.log(lags), np.log(tau), 1)
    return float(slope)


def _load_vix3m_series(index: pd.DatetimeIndex) -> pd.Series:
    """VIX3M cash index; CBOE parquet preferred, else yfinance ^VIX3M."""
    if CBOE_PANEL.is_file():
        ct = pd.read_parquet(CBOE_PANEL)
        ct.index = pd.to_datetime(ct.index).tz_localize(None)
        if "vix3m" in ct.columns:
            s = ct["vix3m"].astype(float)
            return s.reindex(index).ffill().bfill()

    try:
        import yfinance as yf
    except ImportError:
        return pd.Series(np.nan, index=index, name="vix3m")

    start = (index.min() - pd.Timedelta(days=30)).strftime("%Y-%m-%d")
    end = (index.max() + pd.Timedelta(days=5)).strftime("%Y-%m-%d")
    for sym in ("^VIX3M", "VIX3M"):
        raw = yf.download(sym, start=start, end=end, progress=False, auto_adjust=False, threads=False)
        if raw is None or raw.empty:
            continue
        if isinstance(raw.columns, pd.MultiIndex):
            close = raw["Close"].iloc[:, 0]
        else:
            close = raw["Close"]
        close.index = pd.to_datetime(close.index).tz_localize(None)
        return close.rename("vix3m").reindex(index).ffill().bfill()
    return pd.Series(np.nan, index=index, name="vix3m")


@dataclass(frozen=True)
class FeatureConfig:
    rv_window: int = RV_WINDOW
    iv_target_dte: int = IV_TARGET_DTE
    hurst_window: int = HURST_ROLLING
    spot_vol_window: int = SPOT_VOL_WINDOW
    ma_window: int = MA_WINDOW
    skew_put_delta: float = SKEW_PUT_DELTA
    skew_call_delta: float = SKEW_CALL_DELTA
    vrp_median_min_periods: int = 252


def build_regime_features(
    spy_df: pd.DataFrame,
    days: list[pd.Timestamp],
    *,
    iv30_by_day: dict[pd.Timestamp, float | None] | None = None,
    skew25_by_day: dict[pd.Timestamp, float | None] | None = None,
    config: FeatureConfig | None = None,
) -> pd.DataFrame:
    """
    Build continuous + discretized features aligned to ``days``.

    ``iv30_by_day`` / ``skew25_by_day`` optional Theta-derived IV; when missing,
    IV is proxied from VIX/100 and skew is NaN (skew state defaults to ``normal``).
    """
    cfg = config or FeatureConfig()
    o = spy_df.copy()
    o.index = pd.to_datetime(o.index).tz_localize(None)
    o["ret_1"] = o["close"].pct_change()
    o[f"rv{cfg.rv_window}"] = o["ret_1"].rolling(cfg.rv_window).std() * math.sqrt(252.0)
    o[f"sma_{cfg.ma_window}"] = o["close"].rolling(cfg.ma_window, min_periods=1).mean()
    o["ma_dist"] = o["close"] / o[f"sma_{cfg.ma_window}"] - 1.0

    log_px = np.log(o["close"].astype(float))
    o["hurst"] = log_px.rolling(cfg.hurst_window, min_periods=cfg.hurst_window // 2).apply(
        calculate_hurst_exponent,
        raw=False,
    )

    if "vix_close" not in o.columns:
        raise KeyError("spy_df must include vix_close")
    o["spy_ret"] = o["ret_1"]
    o["vix_ret"] = o["vix_close"].pct_change()
    o["spot_vol_corr"] = o["spy_ret"].rolling(cfg.spot_vol_window).corr(o["vix_ret"])

    vix3m = _load_vix3m_series(o.index)
    o["vix3m"] = vix3m
    o["term_ratio"] = o["vix_close"] / o["vix3m"]

    idx = [pd.Timestamp(d).normalize() for d in days]
    panel = o.reindex(idx).ffill()

    iv_map = iv30_by_day or {}
    skew_map = skew25_by_day or {}
    iv_vals: list[float | None] = []
    skew_vals: list[float | None] = []
    for d in idx:
        iv_vals.append(iv_map.get(d))
        skew_vals.append(skew_map.get(d))

    panel["iv30"] = iv_vals
    panel["skew25"] = skew_vals
    # VIX/100 proxy when chain IV unavailable
    proxy_iv = panel["vix_close"] / 100.0
    panel["iv30"] = panel["iv30"].where(panel["iv30"].notna(), proxy_iv)
    panel["vrp30"] = panel["iv30"] - panel[f"rv{cfg.rv_window}"]
    panel["vrp_median"] = panel["vrp30"].expanding(min_periods=cfg.vrp_median_min_periods).median()

    rv_col = f"rv{cfg.rv_window}"
    state_rows = [discretize_row(row, rv_col=rv_col).to_dict() for _, row in panel.iterrows()]
    states = pd.DataFrame(state_rows, index=panel.index)
    return pd.concat([panel, states], axis=1)


def discretize_row(row: pd.Series, *, rv_col: str = f"rv{RV_WINDOW}") -> pd.Series:
    return pd.Series(
        {
            "vrp_state": _discretize_vrp(row.get("vrp30"), row.get("vrp_median")),
            "hurst_state": _discretize_hurst(row.get("hurst")),
            "skew_state": _discretize_skew(row.get("skew25")),
            "term_state": _discretize_term(row.get("term_ratio")),
            "spot_vol_state": _discretize_spot_vol(row.get("spot_vol_corr")),
            "ma_dist_state": _discretize_ma_dist(row.get("ma_dist")),
        }
    )


def discretize_features(
    vrp: float | None,
    hurst: float | None,
    skew: float | None,
    term_ratio: float | None,
    spot_vol_corr: float | None,
    ma_dist: float | None,
    *,
    vrp_median: float | None = None,
) -> FeatureState:
    return (
        _discretize_vrp(vrp, vrp_median),
        _discretize_hurst(hurst),
        _discretize_skew(skew),
        _discretize_term(term_ratio),
        _discretize_spot_vol(spot_vol_corr),
        _discretize_ma_dist(ma_dist),
    )


def _discretize_vrp(vrp: float | None, median: float | None) -> VrpState:
    if vrp is None or not math.isfinite(float(vrp)):
        return "normal"
    v = float(vrp)
    if v < 0.0:
        return "low"
    med = float(median) if median is not None and math.isfinite(float(median)) else 0.0
    return "high" if v > med else "normal"


def _discretize_hurst(h: float | None) -> HurstState:
    if h is None or not math.isfinite(float(h)):
        return "random"
    v = float(h)
    if v < 0.5:
        return "anti"
    if v > 0.5:
        return "persistent"
    return "random"


def _discretize_skew(skew: float | None) -> SkewState:
    if skew is None or not math.isfinite(float(skew)):
        return "normal"
    s = float(skew)
    if s < SKEW_FLAT_HI:
        return "flat"
    if s > SKEW_NORMAL_HI:
        return "steep"
    return "normal"


def _discretize_term(ratio: float | None) -> TermState:
    if ratio is None or not math.isfinite(float(ratio)):
        return "contango"
    return "backwardation" if float(ratio) > 1.0 else "contango"


def _discretize_spot_vol(corr: float | None) -> SpotVolState:
    if corr is None or not math.isfinite(float(corr)):
        return "negative"
    return "positive" if float(corr) >= 0.0 else "negative"


def _discretize_ma_dist(dist: float | None) -> MaDistState:
    if dist is None or not math.isfinite(float(dist)):
        return "neutral"
    d = float(dist)
    if d < -MA_DIST_BAND:
        return "below"
    if d > MA_DIST_BAND:
        return "above"
    return "neutral"
