"""
Rule-based QuantifiedStrategies-style ETF sleeves for ranking / research.

Each function returns a daily return series aligned to ``spy.index`` (cash = 0).
"""

from __future__ import annotations

import numpy as np
import pandas as pd

from RenTech.strategy_stack.run_qs_top_ideas_backtest import (
    _ibs,
    _mom_12m,
    _mr_state_machine,
    _overnight_one_night,
    _rsi,
    _spy_close_to_close,
    dual_momentum_spy_tlt,
    mr_ibs_sma200,
    mr_rsi2_sma200,
    overnight_3down,
    overnight_5d_low,
    overnight_close_to_open,
    rotation_spy_tlt_gld,
    seasonality_opex_week,
    seasonality_turn_of_month,
    seasonality_turnaround_tuesday,
)


def _c2c_on_mask(spy: pd.DataFrame, mask: pd.Series) -> pd.Series:
    return _spy_close_to_close(spy).where(mask, 0.0)


def _month_start_pick(
    idx: pd.DatetimeIndex,
    scores: dict[str, pd.Series],
    rets: dict[str, pd.Series],
    *,
    cash_if_all_neg: bool = False,
) -> pd.Series:
    month = pd.Series(idx, index=idx).dt.to_period("M")
    choice = pd.Series("cash", index=idx, dtype=object)
    for _, dates in month.groupby(month):
        d0 = dates.index[0]
        vals = {k: float(scores[k].loc[d0]) for k in scores}
        if not all(np.isfinite(v) for v in vals.values()):
            pick = list(scores.keys())[0]
        elif cash_if_all_neg and max(vals.values()) <= 0:
            pick = "cash"
        else:
            pick = max(vals, key=vals.get)  # type: ignore[arg-type]
        for d in dates.index:
            choice.loc[d] = pick
    out = pd.Series(0.0, index=idx)
    for d in idx:
        c = choice.loc[d]
        if c != "cash":
            out.loc[d] = float(rets[c].loc[d])
    return out


def _third_friday(year: int, month: int) -> pd.Timestamp:
    d = pd.Timestamp(year=year, month=month, day=1)
    fridays = pd.date_range(d, d + pd.offsets.MonthEnd(0), freq="W-FRI")
    return fridays[2] if len(fridays) >= 3 else fridays[-1]


# --- Seasonality ---


def seasonality_first_day_month(spy: pd.DataFrame, **_kw) -> pd.Series:
    m = pd.Series(spy.index, index=spy.index).dt.to_period("M")
    first = m != m.shift(1)
    return _c2c_on_mask(spy, first)


def seasonality_last_day_month(spy: pd.DataFrame, **_kw) -> pd.Series:
    m = pd.Series(spy.index, index=spy.index).dt.to_period("M")
    last = m != m.shift(-1)
    return _c2c_on_mask(spy, last)


def seasonality_opex_friday(spy: pd.DataFrame, **_kw) -> pd.Series:
    days: set[pd.Timestamp] = set()
    for per in spy.index.to_period("M").unique():
        days.add(_third_friday(per.year, per.month).normalize())
    mask = pd.Series([d.normalize() in days for d in spy.index], index=spy.index)
    return _c2c_on_mask(spy, mask)


def seasonality_sell_in_may(spy: pd.DataFrame, **_kw) -> pd.Series:
    """Long Nov–Apr (Hirsch best six months)."""
    m = spy.index.month
    mask = (m >= 11) | (m <= 4)
    return _c2c_on_mask(spy, mask)


def seasonality_turn_of_year(spy: pd.DataFrame, **_kw) -> pd.Series:
    """Last 5 trading days of Dec + first 2 of Jan."""
    idx = spy.index
    mask = pd.Series(False, index=idx)
    for y in sorted(set(idx.year)):
        dec = idx[(idx.year == y) & (idx.month == 12)]
        jan = idx[(idx.year == y) & (idx.month == 1)]
        if len(dec) >= 5:
            for d in dec[-5:]:
                mask.loc[d] = True
        if len(jan) >= 2:
            for d in jan[:2]:
                mask.loc[d] = True
    return _c2c_on_mask(spy, mask)


def seasonality_monday_reversal(spy: pd.DataFrame, **_kw) -> pd.Series:
    """Long Monday when prior Friday was down."""
    c2c = _spy_close_to_close(spy)
    sig = pd.Series(False, index=spy.index)
    for i in range(1, len(spy)):
        if spy.index[i].dayofweek == 0 and spy.index[i - 1].dayofweek == 4:
            if c2c.iloc[i - 1] < 0:
                sig.iloc[i] = True
    return c2c.where(sig, 0.0)


def seasonality_wednesday(spy: pd.DataFrame, **_kw) -> pd.Series:
    return _c2c_on_mask(spy, spy.index.dayofweek == 2)


def seasonality_friday(spy: pd.DataFrame, **_kw) -> pd.Series:
    return _c2c_on_mask(spy, spy.index.dayofweek == 4)


# --- Overnight ---


def overnight_10d_low(spy: pd.DataFrame, **_kw) -> pd.Series:
    m = spy["close"] <= spy["close"].rolling(10, min_periods=10).min()
    return _overnight_one_night(spy, m)


def overnight_2down(spy: pd.DataFrame, **_kw) -> pd.Series:
    c2c = _spy_close_to_close(spy)
    m = (c2c < 0) & (c2c.shift(1) < 0)
    return _overnight_one_night(spy, m)


def overnight_large_down(spy: pd.DataFrame, **_kw) -> pd.Series:
    c2c = _spy_close_to_close(spy)
    return _overnight_one_night(spy, c2c < -0.01)


def overnight_weekend(spy: pd.DataFrame, **_kw) -> pd.Series:
    """Fri close → Mon open."""
    oc = spy["open"] / spy["close"].shift(1) - 1.0
    out = pd.Series(0.0, index=spy.index)
    for i in range(1, len(spy)):
        if spy.index[i].dayofweek == 0 and spy.index[i - 1].dayofweek == 4:
            out.iloc[i] = float(oc.iloc[i]) if np.isfinite(oc.iloc[i]) else 0.0
    return out


def overnight_gap_down(spy: pd.DataFrame, **_kw) -> pd.Series:
    gap = spy["open"] / spy["close"].shift(1) - 1.0
    return _overnight_one_night(spy, gap < -0.003)


# --- Mean reversion ---


def mr_rsi2_sma200_loose(spy: pd.DataFrame, **_kw) -> pd.Series:
    rsi2 = _rsi(spy["close"], 2)
    sma200 = spy["close"].rolling(200, min_periods=100).mean()
    entry = (rsi2 < 15) & (spy["close"] > sma200)
    return _mr_state_machine(spy, entry, rsi2 > 70, max_hold=5)


def mr_ibs_no_filter(spy: pd.DataFrame, **_kw) -> pd.Series:
    ibs = _ibs(spy)
    entry = ibs < 0.2
    return _mr_state_machine(spy, entry, ibs > 0.8, max_hold=3)


def mr_5day_low_sma200(spy: pd.DataFrame, **_kw) -> pd.Series:
    sma200 = spy["close"].rolling(200, min_periods=100).mean()
    low5 = spy["close"] <= spy["close"].rolling(5, min_periods=5).min()
    entry = low5 & (spy["close"] > sma200)
    exit_sig = spy["close"] >= spy["close"].rolling(5, min_periods=5).max()
    return _mr_state_machine(spy, entry, exit_sig, max_hold=5)


def mr_2day_low_sma200(spy: pd.DataFrame, **_kw) -> pd.Series:
    sma200 = spy["close"].rolling(200, min_periods=100).mean()
    c2c = _spy_close_to_close(spy)
    entry = (c2c < 0) & (c2c.shift(1) < 0) & (spy["close"] > sma200)
    return _mr_state_machine(spy, entry, c2c > 0, max_hold=3)


def mr_rsi5_30_50(spy: pd.DataFrame, **_kw) -> pd.Series:
    rsi5 = _rsi(spy["close"], 5)
    sma200 = spy["close"].rolling(200, min_periods=100).mean()
    entry = (rsi5 < 30) & (spy["close"] > sma200)
    return _mr_state_machine(spy, entry, rsi5 > 50, max_hold=5)


def mr_stoch_sma200(spy: pd.DataFrame, **_kw) -> pd.Series:
    lo = spy["low"].rolling(14, min_periods=14).min()
    hi = spy["high"].rolling(14, min_periods=14).max()
    k = 100 * (spy["close"] - lo) / (hi - lo).replace(0, np.nan)
    sma200 = spy["close"].rolling(200, min_periods=100).mean()
    entry = (k < 20) & (spy["close"] > sma200)
    return _mr_state_machine(spy, entry, k > 80, max_hold=5)


def mr_williams_r_sma200(spy: pd.DataFrame, **_kw) -> pd.Series:
    lo = spy["low"].rolling(14, min_periods=14).min()
    hi = spy["high"].rolling(14, min_periods=14).max()
    wr = -100 * (hi - spy["close"]) / (hi - lo).replace(0, np.nan)
    sma200 = spy["close"].rolling(200, min_periods=100).mean()
    entry = (wr < -80) & (spy["close"] > sma200)
    return _mr_state_machine(spy, entry, wr > -20, max_hold=5)


def mr_connors_3up_short(spy: pd.DataFrame, **_kw) -> pd.Series:
    """Fade 3 up days: short next day (negative SPY return)."""
    c2c = _spy_close_to_close(spy)
    sig = (c2c > 0) & (c2c.shift(1) > 0) & (c2c.shift(2) > 0)
    return (-c2c).where(sig.shift(1).fillna(False), 0.0)


def mr_large_drop_bounce(spy: pd.DataFrame, **_kw) -> pd.Series:
    c2c = _spy_close_to_close(spy)
    entry = c2c < -0.02
    return _mr_state_machine(spy, entry, c2c > 0, max_hold=2)


def mr_nr7(spy: pd.DataFrame, **_kw) -> pd.Series:
    rng = spy["high"] - spy["low"]
    nr7 = rng <= rng.rolling(7, min_periods=7).min()
    return _mr_state_machine(spy, nr7.shift(1).fillna(False), rng > rng.rolling(7).mean(), max_hold=3)


# --- Momentum / rotation ---


def momentum_12m_spy_cash(spy: pd.DataFrame, **_kw) -> pd.Series:
    mom = _mom_12m(spy["close"])
    mask = mom > 0
    return _c2c_on_mask(spy, mask)


def rotation_spy_tlt(spy: pd.DataFrame, tlt: pd.DataFrame, **_kw) -> pd.Series:
    idx = spy.index.intersection(tlt.index)
    panels = {"spy": spy.loc[idx, "close"], "tlt": tlt.loc[idx, "close"]}
    rets = {k: v.pct_change().fillna(0.0) for k, v in panels.items()}
    moms = {k: _mom_12m(v) for k, v in panels.items()}
    return _month_start_pick(idx, moms, rets)


def faber_200sma_spy_cash(spy: pd.DataFrame, **_kw) -> pd.Series:
    sma = spy["close"].rolling(200, min_periods=100).mean()
    return _c2c_on_mask(spy, spy["close"] > sma)


def faber_200sma_spy_tlt(spy: pd.DataFrame, tlt: pd.DataFrame, **_kw) -> pd.Series:
    idx = spy.index.intersection(tlt.index)
    sma = spy.loc[idx, "close"].rolling(200, min_periods=100).mean()
    spy_r = spy.loc[idx, "close"].pct_change().fillna(0.0)
    tlt_r = tlt.loc[idx, "close"].pct_change().fillna(0.0)
    long_spy = spy.loc[idx, "close"] > sma
    out = pd.Series(0.0, index=idx)
    out.loc[long_spy] = spy_r.loc[long_spy]
    out.loc[~long_spy] = tlt_r.loc[~long_spy]
    return out


def golden_cross_50_200(spy: pd.DataFrame, **_kw) -> pd.Series:
    s50 = spy["close"].rolling(50, min_periods=40).mean()
    s200 = spy["close"].rolling(200, min_periods=100).mean()
    return _c2c_on_mask(spy, s50 > s200)


def rotation_defensive_xlp_xlu(
    spy: pd.DataFrame, xlp: pd.DataFrame, xlu: pd.DataFrame, **_kw
) -> pd.Series:
    idx = spy.index.intersection(xlp.index).intersection(xlu.index)
    panels = {"spy": spy.loc[idx, "close"], "xlp": xlp.loc[idx, "close"], "xlu": xlu.loc[idx, "close"]}
    rets = {k: v.pct_change().fillna(0.0) for k, v in panels.items()}
    moms = {k: _mom_12m(v) for k, v in panels.items()}
    return _month_start_pick(idx, moms, rets)


def momentum_6m_spy_cash(spy: pd.DataFrame, **_kw) -> pd.Series:
    mom = spy["close"] / spy["close"].shift(126) - 1.0
    return _c2c_on_mask(spy, mom > 0)


# --- Volatility / VIX ---


def vol_vix_spike_fade(spy: pd.DataFrame, vix: pd.DataFrame, **_kw) -> pd.Series:
    """After VIX jumps >15% in a day, long SPY next session."""
    vc = vix["close"].reindex(spy.index).ffill()
    vchg = vc.pct_change()
    entry = vchg > 0.15
    return _mr_state_machine(spy, entry, pd.Series(False, index=spy.index), max_hold=2)


def vol_vix_high_cash(spy: pd.DataFrame, vix: pd.DataFrame, **_kw) -> pd.Series:
    vc = vix["close"].reindex(spy.index).ffill()
    return _c2c_on_mask(spy, vc < 25)


def vol_vix_low_risk_on(spy: pd.DataFrame, vix: pd.DataFrame, **_kw) -> pd.Series:
    vc = vix["close"].reindex(spy.index).ffill()
    return _c2c_on_mask(spy, vc < 15)


def vol_bollinger_squeeze(spy: pd.DataFrame, **_kw) -> pd.Series:
    c = spy["close"]
    mid = c.rolling(20, min_periods=20).mean()
    std = c.rolling(20, min_periods=20).std()
    width = (2 * std) / mid
    squeeze = width < width.rolling(126, min_periods=60).quantile(0.2)
    breakout = c > mid + 2 * std
    entry = squeeze.shift(1).fillna(False) & breakout
    return _mr_state_machine(spy, entry, c < mid, max_hold=10)


# --- Price action ---


def pa_inside_day_breakout(spy: pd.DataFrame, **_kw) -> pd.Series:
    inside = (spy["high"] < spy["high"].shift(1)) & (spy["low"] > spy["low"].shift(1))
    brk = spy["close"] > spy["high"].shift(1)
    entry = inside.shift(1).fillna(False) & brk
    return _mr_state_machine(spy, entry, spy["close"] < spy["low"].shift(1), max_hold=5)


def pa_gap_up_fade(spy: pd.DataFrame, **_kw) -> pd.Series:
    gap = spy["open"] / spy["close"].shift(1) - 1.0
    return (-_spy_close_to_close(spy)).where(gap > 0.005, 0.0)


def pa_52w_high_momentum(spy: pd.DataFrame, **_kw) -> pd.Series:
    hi = spy["close"].rolling(252, min_periods=200).max()
    return _c2c_on_mask(spy, spy["close"] >= hi * 0.98)


# --- Composites ---


def combo_turn_tue_3down_ibs(spy: pd.DataFrame, **_kw) -> pd.Series:
    a = seasonality_turnaround_tuesday(spy)
    b = overnight_3down(spy)
    c = mr_ibs_sma200(spy)
    return (a + b + c) / 3.0


def combo_seasonality_stack(spy: pd.DataFrame, **_kw) -> pd.Series:
    a = seasonality_turnaround_tuesday(spy)
    b = seasonality_turn_of_month(spy)
    c = seasonality_monday_reversal(spy)
    active = (a != 0) | (b != 0) | (c != 0)
    c2c = _spy_close_to_close(spy)
    return c2c.where(active, 0.0)


# Registry: 50 systematic sleeves
SYSTEMATIC_STRATEGIES: dict[str, object] = {
  # seasonality (12)
    "S01_turnaround_tuesday": seasonality_turnaround_tuesday,
    "S02_turn_of_month": seasonality_turn_of_month,
    "S03_first_day_month": seasonality_first_day_month,
    "S04_last_day_month": seasonality_last_day_month,
    "S05_opex_week": seasonality_opex_week,
    "S06_opex_friday": seasonality_opex_friday,
    "S07_sell_in_may": seasonality_sell_in_may,
    "S08_turn_of_year": seasonality_turn_of_year,
    "S09_monday_reversal": seasonality_monday_reversal,
    "S10_wednesday": seasonality_wednesday,
    "S11_friday": seasonality_friday,
    "S12_seasonality_combo": combo_seasonality_stack,
    # overnight (8)
    "S13_overnight_close_open": overnight_close_to_open,
    "S14_overnight_3down": overnight_3down,
    "S15_overnight_2down": overnight_2down,
    "S16_overnight_5d_low": overnight_5d_low,
    "S17_overnight_10d_low": overnight_10d_low,
    "S18_overnight_large_down": overnight_large_down,
    "S19_overnight_weekend": overnight_weekend,
    "S20_overnight_gap_down": overnight_gap_down,
    # mean reversion (12)
    "S21_mr_rsi2_sma200": mr_rsi2_sma200,
    "S22_mr_rsi2_loose": mr_rsi2_sma200_loose,
    "S23_mr_ibs_sma200": mr_ibs_sma200,
    "S24_mr_ibs_nofilter": mr_ibs_no_filter,
    "S25_mr_5d_low_sma200": mr_5day_low_sma200,
    "S26_mr_2d_down_sma200": mr_2day_low_sma200,
    "S27_mr_rsi5_30_50": mr_rsi5_30_50,
    "S28_mr_stoch_sma200": mr_stoch_sma200,
    "S29_mr_williams_r": mr_williams_r_sma200,
    "S30_mr_connors_3up_fade": mr_connors_3up_short,
    "S31_mr_large_drop": mr_large_drop_bounce,
    "S32_mr_nr7": mr_nr7,
    # momentum / rotation (10)
    "S33_dual_momentum_spy_tlt": dual_momentum_spy_tlt,
    "S34_rotation_spy_tlt_gld": rotation_spy_tlt_gld,
    "S35_rotation_spy_tlt": rotation_spy_tlt,
    "S36_momentum_12m_cash": momentum_12m_spy_cash,
    "S37_momentum_6m_cash": momentum_6m_spy_cash,
    "S38_faber_200sma_cash": faber_200sma_spy_cash,
    "S39_faber_200sma_tlt": faber_200sma_spy_tlt,
    "S40_golden_cross_50_200": golden_cross_50_200,
    "S41_rotation_xlp_xlu_spy": rotation_defensive_xlp_xlu,
    "S42_52w_high_momentum": pa_52w_high_momentum,
    # volatility (5)
    "S43_vix_spike_fade": vol_vix_spike_fade,
    "S44_vix_high_cash": vol_vix_high_cash,
    "S45_vix_low_risk_on": vol_vix_low_risk_on,
    "S46_bollinger_squeeze": vol_bollinger_squeeze,
    "S47_gap_up_fade": pa_gap_up_fade,
    # composites / PA (3)
    "S48_inside_day_breakout": pa_inside_day_breakout,
    "S49_combo_tue_3down_ibs": combo_turn_tue_3down_ibs,
    "S50_qqq_rsi2_sma200": None,  # placeholder filled at runtime
}


def qqq_rsi2_sma200(spy: pd.DataFrame, qqq: pd.DataFrame, **_kw) -> pd.Series:
    rsi2 = _rsi(qqq["close"].reindex(spy.index).ffill(), 2)
    sma200 = qqq["close"].reindex(spy.index).ffill().rolling(200, min_periods=100).mean()
    qqq_r = qqq["close"].reindex(spy.index).ffill().pct_change().fillna(0.0)
    entry = (rsi2 < 10) & (qqq["close"].reindex(spy.index).ffill() > sma200)
    pos = np.zeros(len(spy), dtype=bool)
    hold = 0
    for i in range(len(spy)):
        prev = pos[i - 1] if i else False
        if prev:
            hold += 1
            if bool(rsi2.iloc[i - 1] > 70) or hold > 5:
                pos[i] = False
                hold = 0
            else:
                pos[i] = True
        elif i > 0 and bool(entry.iloc[i - 1]):
            pos[i] = True
            hold = 1
    return qqq_r.where(pd.Series(pos, index=spy.index), 0.0)


SYSTEMATIC_STRATEGIES["S50_qqq_rsi2_sma200"] = qqq_rsi2_sma200

# Stock-only book: four non-redundant SPY diversifiers (no MR overlap with CM dip).
QS_ACTIONABLE_4: tuple[str, ...] = (
    "S01_turnaround_tuesday",
    "S14_overnight_3down",
    "S03_first_day_month",
    "S17_overnight_10d_low",
)

# Full options-book diversifier set (non-redundant vs tactical AW / VRP / sector mom).
QS_ACTIONABLE_7: tuple[str, ...] = (
    "S14_overnight_3down",
    "S49_combo_tue_3down_ibs",
    "S01_turnaround_tuesday",
    "S17_overnight_10d_low",
    "S21_mr_rsi2_sma200",
    "S23_mr_ibs_sma200",
    "S03_first_day_month",
)

QS_ACTIONABLE_PRESETS: dict[str, tuple[str, ...]] = {
    "actionable-4": QS_ACTIONABLE_4,
    "actionable-7": QS_ACTIONABLE_7,
}


def run_systematic_strategy(
    sid: str,
    panels: dict[str, pd.DataFrame],
    idx: pd.DatetimeIndex,
) -> pd.Series:
    """Run one registry sleeve; returns daily returns on *idx*."""
    import inspect

    fn = SYSTEMATIC_STRATEGIES[sid]
    if fn is None:
        raise KeyError(sid)
    sig = inspect.signature(fn)
    kwargs: dict = {}
    for name in ("tlt", "gld", "vix", "xlp", "xlu", "qqq"):
        if name in sig.parameters:
            kwargs[name] = panels[name]
    r = fn(panels["spy"], **kwargs)
    return r.reindex(idx).fillna(0.0)


def actionable_daily_return(
    panels: dict[str, pd.DataFrame],
    idx: pd.DatetimeIndex,
    members: tuple[str, ...],
) -> pd.Series:
    parts = [run_systematic_strategy(sid, panels, idx) for sid in members]
    return pd.concat(parts, axis=1).mean(axis=1)


def actionable_4_daily_return(
    panels: dict[str, pd.DataFrame],
    idx: pd.DatetimeIndex,
) -> pd.Series:
    return actionable_daily_return(panels, idx, QS_ACTIONABLE_4)


def actionable_7_daily_return(
    panels: dict[str, pd.DataFrame],
    idx: pd.DatetimeIndex,
) -> pd.Series:
    return actionable_daily_return(panels, idx, QS_ACTIONABLE_7)
