"""
Standalone **inverse SPY ETF** bear hedge via dual MA-slope + SPY regime gates.

Firmed-up rules (vs legacy inverse-ETF-slope-only):
  * **SPY regime gate** — only hedge when SPY is in a bear state (dual negative
    slope and/or below SMA200).
  * **Sticky SPY exit** — stay hedged until SPY bull resumes, not on inverse-ETF
    bounce whipsaws.
  * **SH default** — −1x instrument avoids 2x/3x decay.

Separate from the S&P 500 top-N long sleeve.
"""

from __future__ import annotations

from dataclasses import dataclass
from typing import Dict, Literal

import numpy as np
import pandas as pd

from RenTech.strategy_stack.ma_slope_engine import (
    ExitMode,
    MaSlopeConfig,
    MaSlopeEngine,
    compute_ma,
    compute_ma_slope,
    compute_slope_rank_score,
    compute_vol_scale,
)
from RenTech.strategy_stack.multi_strategy_manager import _align_panel_frames

INVERSE_SPY_ETFS: dict[str, str] = {
    "SH": "ProShares Short S&P500 (-1x)",
    "SDS": "ProShares UltraShort S&P500 (-2x)",
    "SPXU": "ProShares UltraPro Short S&P500 (-3x)",
}

SelectionMode = Literal["all_active", "top_n"]
SpyRegime = Literal["none", "bear_dual_slope", "below_sma200", "bear_dual_or_sma200"]
SpyExit = Literal[
    "inverse_signal",
    "spy_fast_slope_positive",
    "spy_regime_off",
    "spy_slope_or_sma200",
    "spy_bull_both",
]


@dataclass(frozen=True)
class MaSlopeInverseConfig:
    fast_period: int = 10
    slow_period: int = 50
    fast_lookback: int = 10
    slow_lookback: int = 5
    entry_slope_min: float = 0.0
    price_above_ma: bool = True
    exit_mode: ExitMode = "slope_flip"
    selection: SelectionMode = "top_n"
    top_n: int = 1
    tickers_preferred: tuple[str, ...] = ("SH",)
    # Firmed-up bear hedge
    spy_regime: SpyRegime = "bear_dual_or_sma200"
    spy_sma_window: int = 200
    spy_exit: SpyExit = "spy_slope_or_sma200"
    require_inverse_momentum: bool = True
    sizing_mode: str = "binary"
    target_vol: float = 0.15
    vol_scale_floor: float = 0.5
    vol_scale_cap: float = 1.5
    vol_window: int = 20
    cash_annual_yield: float = 0.04

    def slug(self) -> str:
        inv = "invreq" if self.require_inverse_momentum else "spybear"
        return (
            f"{self.spy_regime}_{self.spy_exit}_{inv}_"
            f"{self.selection}{self.top_n}_{','.join(self.tickers_preferred)}"
        )


@dataclass
class MaSlopeInverseSleeve:
    config: MaSlopeInverseConfig = MaSlopeInverseConfig()

    def _engine_config(self) -> MaSlopeConfig:
        c = self.config
        return MaSlopeConfig(
            ma_type="ema",
            ma_period=c.fast_period,
            slope_lookback=c.fast_lookback,
            slope_method="pct",
            entry_slope_min=c.entry_slope_min,
            price_above_ma=c.price_above_ma,
            direction="long_only",
            exit_mode=c.exit_mode,
            signal_mode="dual_timeframe",
            slow_ma_period=c.slow_period,
            slow_slope_lookback=c.slow_lookback,
            sizing_mode="binary",
        )

    def _compute_spy_regime(self, spy_df: pd.DataFrame) -> tuple[pd.Series, pd.Series, pd.Series]:
        """Return (bear_regime, fast_slope, exit_signal) aligned to spy index."""
        cfg = self.config
        close = spy_df["close"].astype(np.float64)
        fast_ma = compute_ma(close, ma_type="ema", period=cfg.fast_period)
        slow_ma = compute_ma(close, ma_type="ema", period=cfg.slow_period)
        fast_slope = compute_ma_slope(fast_ma, lookback=cfg.fast_lookback)
        slow_slope = compute_ma_slope(slow_ma, lookback=cfg.slow_lookback)
        sma = close.rolling(cfg.spy_sma_window, min_periods=cfg.spy_sma_window).mean()

        bear_dual = (fast_slope < 0) & (slow_slope < 0)
        below_sma = close < sma

        if cfg.spy_regime == "none":
            bear = pd.Series(True, index=close.index)
        elif cfg.spy_regime == "bear_dual_slope":
            bear = bear_dual
        elif cfg.spy_regime == "below_sma200":
            bear = below_sma
        elif cfg.spy_regime == "bear_dual_or_sma200":
            bear = bear_dual | below_sma
        else:
            raise ValueError(f"unknown spy_regime: {cfg.spy_regime!r}")

        if cfg.spy_exit == "spy_fast_slope_positive":
            spy_exit = fast_slope > 0
        elif cfg.spy_exit == "spy_regime_off":
            spy_exit = ~bear
        elif cfg.spy_exit == "spy_slope_or_sma200":
            spy_exit = (fast_slope > 0) | (close > sma)
        elif cfg.spy_exit == "spy_bull_both":
            spy_exit = (fast_slope > 0) & (close > sma)
        elif cfg.spy_exit == "inverse_signal":
            spy_exit = pd.Series(False, index=close.index)
        else:
            raise ValueError(f"unknown spy_exit: {cfg.spy_exit!r}")

        return bear.fillna(False), fast_slope, spy_exit.fillna(False)

    def _inverse_active_panel(self, etf_dict: Dict[str, pd.DataFrame]) -> tuple[pd.DataFrame, pd.DataFrame]:
        cfg = self.config
        eng = MaSlopeEngine(config=self._engine_config())
        active_pan: list[pd.Series] = []
        score_pan: list[pd.Series] = []
        for t, df in sorted(etf_dict.items()):
            if df is None or df.empty:
                continue
            idx = pd.to_datetime(df.index).tz_localize(None)
            dfx = df.copy()
            dfx.index = idx
            dfx = dfx.sort_index()
            frame = eng.transform(dfx)
            scored = compute_slope_rank_score(
                dfx["close"].astype(np.float64),
                fast_period=cfg.fast_period,
                slow_period=cfg.slow_period,
                fast_lookback=cfg.fast_lookback,
                slow_lookback=cfg.slow_lookback,
                entry_slope_min=cfg.entry_slope_min,
                price_above_ma=cfg.price_above_ma,
                rank_metric="dual_product",
            )
            active = frame["target_position"].astype(np.float64) > 0
            score = scored["rank_score"].where(active, np.nan)
            active_pan.append(active.astype(np.float64).rename(t))
            score_pan.append(score.rename(t))
        active_df, _ = _align_panel_frames(active_pan)
        score_df, _ = _align_panel_frames(score_pan)
        return active_df.sort_index(), score_df.sort_index()

    def build_weight_panel(
        self,
        etf_dict: Dict[str, pd.DataFrame],
        spy_df: pd.DataFrame,
    ) -> tuple[pd.DataFrame, pd.DataFrame]:
        """
        Daily weights per inverse ETF using SPY regime state machine.

        Returns (weights, rank_score_panel).
        """
        if not etf_dict:
            raise ValueError("etf_dict cannot be empty")
        cfg = self.config
        active_df, score_df = self._inverse_active_panel(etf_dict)
        bear, _, spy_exit = self._compute_spy_regime(spy_df)

        master = active_df.index.union(bear.index).sort_values()
        active_df = active_df.reindex(master).fillna(0.0)
        score_df = score_df.reindex(master)
        bear = bear.reindex(master).fillna(False)
        spy_exit = spy_exit.reindex(master).fillna(False)

        tickers = list(active_df.columns)
        n_t = len(tickers)
        n_d = len(master)
        out_arr = np.zeros((n_d, n_t), dtype=np.float64)

        # Optional inverse exit row (legacy path)
        inv_exit_any = pd.Series(False, index=master)
        if cfg.spy_exit == "inverse_signal" or cfg.require_inverse_momentum:
            eng = MaSlopeEngine(config=self._engine_config())
            for j, t in enumerate(tickers):
                df = etf_dict[t]
                idx = pd.to_datetime(df.index).tz_localize(None)
                dfx = df.copy()
                dfx.index = idx
                frame = eng.transform(dfx.sort_index())
                inv_exit_any = inv_exit_any | (frame["long_exit"].reindex(master).fillna(False))

        state = 0
        for i in range(n_d):
            if state == 0:
                entry_ok = bool(bear.iloc[i])
                if cfg.require_inverse_momentum:
                    inv_any = bool(active_df.iloc[i].max() > 0)
                    entry_ok = entry_ok and inv_any
                if entry_ok:
                    state = 1
            else:
                if cfg.spy_exit == "inverse_signal":
                    if bool(inv_exit_any.iloc[i]):
                        state = 0
                elif bool(spy_exit.iloc[i]):
                    state = 0
                elif cfg.require_inverse_momentum and not bool(active_df.iloc[i].max() > 0):
                    state = 0

            if state == 0:
                continue

            active_idx = np.nonzero(active_df.iloc[i].to_numpy(dtype=np.float64) > 0)[0]
            if cfg.require_inverse_momentum and active_idx.size == 0:
                continue

            if not cfg.require_inverse_momentum:
                # Hold preferred ticker(s) directly when SPY bear
                chosen_names = [t for t in cfg.tickers_preferred if t in tickers]
                if not chosen_names:
                    chosen_names = tickers[:1]
                for name in chosen_names:
                    out_arr[i, tickers.index(name)] = 1.0 / len(chosen_names)
            elif cfg.selection == "all_active":
                for j in active_idx:
                    out_arr[i, j] = 1.0 / float(len(active_idx))
            else:
                sc_row = score_df.iloc[i].to_numpy(dtype=np.float64)
                sub_sc = sc_row[active_idx]
                valid = np.isfinite(sub_sc) & (sub_sc > 0)
                if valid.any():
                    cand = active_idx[valid]
                    k = min(int(cfg.top_n), cand.size)
                    pick = np.argpartition(-sub_sc[valid], k - 1)[:k]
                    chosen = cand[pick]
                else:
                    chosen = active_idx[: min(int(cfg.top_n), active_idx.size)]
                for j in chosen:
                    out_arr[i, j] = 1.0 / float(len(chosen))

        weights = pd.DataFrame(out_arr, index=master, columns=tickers)
        weights = weights.shift(1).fillna(0.0)
        return weights, score_df

    def generate_returns(
        self,
        etf_dict: Dict[str, pd.DataFrame],
        spy_df: pd.DataFrame,
        *,
        verbose: bool = True,
    ) -> pd.Series:
        ret_pan: list[pd.Series] = []
        for t, df in sorted(etf_dict.items()):
            if df is None or df.empty or "ret" not in df.columns:
                continue
            idx = pd.to_datetime(df.index).tz_localize(None)
            dfx = df.copy()
            dfx.index = idx
            ret_pan.append(dfx["ret"].astype(np.float64).rename(t))

        if not ret_pan:
            raise ValueError("no ETF return series")

        ret_df, _ = _align_panel_frames(ret_pan)
        weights, _ = self.build_weight_panel(etf_dict, spy_df)
        weights = weights.reindex(columns=ret_df.columns).fillna(0.0)
        master = ret_df.index.sort_values()
        ret_df = ret_df.reindex(master).fillna(0.0)
        weights = weights.reindex(master).fillna(0.0)

        daily_rf = float(self.config.cash_annual_yield) / 252.0
        w = weights.to_numpy(dtype=np.float64)
        r = ret_df.to_numpy(dtype=np.float64)
        core = (w * r).sum(axis=1)
        wsum = w.sum(axis=1)

        if self.config.sizing_mode == "vol_scale":
            pre = pd.Series(core + np.maximum(0.0, 1.0 - wsum) * daily_rf, index=master)
            scale = compute_vol_scale(
                pre,
                target_vol=self.config.target_vol,
                window=self.config.vol_window,
                floor=self.config.vol_scale_floor,
                cap=self.config.vol_scale_cap,
            )
            out = pre * scale.shift(1).fillna(1.0)
        else:
            out = core + np.maximum(0.0, 1.0 - wsum) * daily_rf

        if verbose:
            invested = float((w.sum(axis=1) > 0).mean())
            bear, _, _ = self._compute_spy_regime(spy_df)
            bear_frac = float(bear.reindex(master).fillna(False).mean())
            print(
                f"  Bear hedge: regime={self.config.spy_regime} exit={self.config.spy_exit} | "
                f"ETFs={list(etf_dict.keys())} | invested {invested:.1%} | "
                f"SPY bear {bear_frac:.1%}",
                flush=True,
            )

        return pd.Series(out, index=master, name="ma_slope_inverse_ret", dtype=np.float64)

    def generate_position_log(
        self,
        etf_dict: Dict[str, pd.DataFrame],
        spy_df: pd.DataFrame,
    ) -> pd.DataFrame:
        weights, scores = self.build_weight_panel(etf_dict, spy_df)
        bear, _, _ = self._compute_spy_regime(spy_df)
        bear = bear.reindex(weights.index).fillna(False)
        rows: list[dict] = []
        for dt in weights.index:
            for tkr in weights.columns:
                w = float(weights.loc[dt, tkr])
                if w <= 0:
                    continue
                sc = scores.loc[dt, tkr] if tkr in scores.columns else float("nan")
                rows.append(
                    {
                        "date": pd.Timestamp(dt).strftime("%Y-%m-%d"),
                        "ticker": tkr,
                        "weight": w,
                        "rank_score": float(sc) if np.isfinite(sc) else float("nan"),
                        "spy_bear_regime": bool(bear.loc[dt]),
                    }
                )
        return pd.DataFrame(rows)
