"""
Combine macro momentum regime (daily) with micro mean-reversion positions (intraday).

Rule 1 — Regime +1 (bull): only **long** micro positions allowed; shorts forced flat.
Rule 2 — Regime -1 (bear): only **short** micro positions allowed; longs forced flat.
Regime 0 (flat / mixed): **no** positions (conservative default).
"""

from __future__ import annotations

from dataclasses import dataclass

import numpy as np
import pandas as pd


def align_daily_regime_to_dates(
    regime: pd.Series,
    dates: pd.Series,
    *,
    shift_sessions: int = 0,
) -> pd.Series:
    """
    Map each row's calendar ``dates`` to the daily regime value.

    Parameters
    ----------
    regime
        Daily regime indexed by datetime (any time component; normalized to date).
    dates
        Per-row calendar date (e.g. ``trade_date`` from DataLoader.align_to_trading_days).
    shift_sessions
        If > 0, use regime from *shift_sessions* calendar days **before** ``dates``
        (approximation of prior sessions; not holiday-aware).
    """
    r = regime.copy()
    r.index = pd.to_datetime(r.index).normalize()
    if r.index.tz is not None:
        r.index = r.index.tz_localize(None)

    d = pd.to_datetime(dates).dt.normalize()
    if getattr(d.dt, "tz", None) is not None:
        d = d.dt.tz_localize(None)

    if shift_sessions > 0:
        d = d - pd.Timedelta(days=int(shift_sessions))

    union_idx = r.index.union(pd.Index(d.dropna().unique())).sort_values()
    rd = r.reindex(union_idx).ffill().fillna(0).astype(np.int8)
    return d.map(rd).fillna(0).astype(np.int8)


def apply_momentum_gate(
    micro_position: pd.Series,
    regime_on_bar: pd.Series,
) -> pd.Series:
    """
    Apply bull/bear gating to micro {-1,0,1} positions.

    Returns orchestrated position series (same index).
    """
    m = micro_position.astype(np.int8)
    reg = regime_on_bar.astype(np.int8)
    out = np.zeros(len(m), dtype=np.int8)

    for i in range(len(m)):
        ri = int(reg.iloc[i])
        pi = int(m.iloc[i])
        if ri == 1:
            out[i] = pi if pi >= 0 else 0
        elif ri == -1:
            out[i] = pi if pi <= 0 else 0
        else:
            out[i] = 0

    return pd.Series(out, index=m.index, dtype=np.int8)


@dataclass
class StrategyOrchestrator:
    """
    Parameters
    ----------
    regime_shift_sessions
        Forward each intraday row to regime from N sessions earlier (0 = same calendar date).
    """

    regime_shift_sessions: int = 0

    def merge_and_gate(
        self,
        intraday: pd.DataFrame,
        daily_regime: pd.Series,
        *,
        trade_date_col: str = "trade_date",
        micro_col: str = "micro_position",
    ) -> pd.DataFrame:
        """
        Add columns ``regime``, ``orchestrated_position``.

        Requires ``intraday[trade_date_col]`` and ``intraday[micro_col]``.
        """
        if trade_date_col not in intraday.columns:
            raise KeyError(f"Missing {trade_date_col}")
        if micro_col not in intraday.columns:
            raise KeyError(f"Missing {micro_col}")

        out = intraday.copy()
        out["regime"] = align_daily_regime_to_dates(
            daily_regime,
            out[trade_date_col],
            shift_sessions=self.regime_shift_sessions,
        )
        out["orchestrated_position"] = apply_momentum_gate(out[micro_col], out["regime"])
        return out


if __name__ == "__main__":
    idx = pd.date_range("2024-06-01", periods=10, freq="h")
    intra = pd.DataFrame(
        {
            "trade_date": pd.to_datetime(["2024-06-03"] * 10),
            "micro_position": [0, 1, 1, -1, 0, 0, -1, -1, 1, 0],
        },
        index=idx,
    )
    reg = pd.Series([1], index=pd.to_datetime(["2024-06-03"]))
    orch = StrategyOrchestrator()
    merged = orch.merge_and_gate(intra, reg)
    print(merged[["micro_position", "regime", "orchestrated_position"]])
