"""
5-minute Opening Range Breakout on Stocks in Play (Zarattini, Barbon, Aziz).

Paper: "A Profitable Day Trading Strategy For The U.S. Equity Market" (SSRN 4729284).

Universe filters (daily):
  - open price > $5
  - 14-day average daily volume >= 1M shares
  - 14-day ATR > $0.50
  - opening-range (first 5m) relative volume >= 100%
  - top N (default 20) by relative volume that session

Entry (after first 5m bar):
  - bullish OR (close > open): stop buy at OR high
  - bearish OR (close < open): stop sell at OR low
  - doji (open == close): no trade

Exit:
  - stop loss at 10% of prior-day ATR from entry
  - otherwise flat at session close (MOC)

Sizing:
  - 1% of equity at risk if stop is hit
  - max 4x gross leverage vs equity at day open
  - commission $0.0035/share (IB tiered default in paper)
"""

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.cm_intraday_dip import (
    TradeRow,
    _session_bar_index,
    _session_key,
    _slip,
    _wilder_atr,
    portfolio_metrics,
    trade_stats,
)

Side = Literal["long", "short"]


@dataclass(frozen=True)
class OrbZarattiniConfig:
    min_price: float = 5.0
    min_avg_vol_14: float = 1_000_000.0
    min_atr_14: float = 0.50
    min_rel_vol: float = 1.0
    top_n: int = 20
    or_lookback: int = 14
    atr_period: int = 14
    # Opening-range length in minutes; panel bar size must divide this evenly.
    or_minutes: int = 5
    bar_minutes: int = 5
    stop_atr_frac: float = 0.10
    risk_per_trade: float = 0.01
    max_leverage: float = 4.0
    commission_per_share: float = 0.0035
    slippage_bps: float = 3.0
    starting_capital: float = 25_000.0
    # If True, size every trade off starting_capital (portfolio sleeve unit).
    # If False, size off compounding equity (paper replication).
    fixed_sizing: bool = False
    # OR quality: range as fraction of prior-day ATR (0 = filter off).
    min_or_range_atr: float = 0.0
    max_or_range_atr: float = 0.0
    # |OR close-open| / OR range; 0 = off. Strong directional OR candle.
    min_body_frac: float = 0.0
    # Last session bar index allowed for entry (0 = no cap). On 5m RTH, bar 18 ≈ 11:00 ET.
    max_entry_bar: int = 0
    # Rank-weight: size shares by rel_vol rank within the day's picks (False = equal risk).
    rvol_rank_size: bool = False


def zarattini_config() -> OrbZarattiniConfig:
    return OrbZarattiniConfig()


def _or_bar_count(cfg: OrbZarattiniConfig) -> int:
    if cfg.or_minutes <= 0 or cfg.bar_minutes <= 0:
        raise ValueError("or_minutes and bar_minutes must be > 0")
    if cfg.or_minutes % cfg.bar_minutes != 0:
        raise ValueError(f"or_minutes={cfg.or_minutes} must be divisible by bar_minutes={cfg.bar_minutes}")
    return int(cfg.or_minutes // cfg.bar_minutes)


@dataclass
class OrScreenRow:
    session_date: pd.Timestamp
    symbol: str
    rel_vol: float
    or_open: float
    or_high: float
    or_low: float
    or_close: float
    or_volume: float
    prior_atr: float
    avg_vol_14: float
    side: Side | None


def _extract_or_rows(sym: str, intra: pd.DataFrame, daily: pd.DataFrame, cfg: OrbZarattiniConfig) -> list[OrScreenRow]:
    if intra.empty or daily.empty:
        return []

    or_bars = _or_bar_count(cfg)
    idx = pd.to_datetime(intra.index).tz_localize(None)
    o = intra["open"].astype(np.float64)
    h = intra["high"].astype(np.float64)
    l = intra["low"].astype(np.float64)
    c = intra["close"].astype(np.float64)
    v = intra["volume"].astype(np.float64) if "volume" in intra.columns else pd.Series(0.0, index=idx)

    d_close = daily["close"].astype(np.float64)
    d_high = daily["high"].astype(np.float64)
    d_low = daily["low"].astype(np.float64)
    d_vol = daily["volume"].astype(np.float64) if "volume" in daily.columns else pd.Series(np.nan, index=daily.index)
    d_atr = _wilder_atr(d_high, d_low, d_close, cfg.atr_period)
    avg_vol_14 = d_vol.rolling(cfg.or_lookback, min_periods=cfg.or_lookback).mean()

    ses = _session_key(idx)
    bar_n = _session_bar_index(idx)
    or_mask = bar_n < or_bars
    if not or_mask.any():
        return []

    tmp = pd.DataFrame(
        {
            "session_date": ses.values,
            "bar_n": bar_n.values,
            "open": o.values,
            "high": h.values,
            "low": l.values,
            "close": c.values,
            "volume": v.values,
        },
        index=idx,
    )
    or_tmp = tmp.loc[or_mask]
    # Require a full opening-range window (all or_bars present).
    counts = or_tmp.groupby("session_date").size()
    complete = counts[counts >= or_bars].index
    or_tmp = or_tmp[or_tmp["session_date"].isin(complete)]
    if or_tmp.empty:
        return []

    g = or_tmp.groupby("session_date", sort=True)
    f = pd.DataFrame(
        {
            "or_open": g["open"].first(),
            "or_high": g["high"].max(),
            "or_low": g["low"].min(),
            "or_close": g["close"].last(),
            "or_volume": g["volume"].sum(),
        }
    ).reset_index()
    f["session_date"] = pd.to_datetime(f["session_date"])
    f = f.merge(
        avg_vol_14.rename("avg_vol_14"),
        left_on="session_date",
        right_index=True,
        how="left",
    )
    f = f.merge(
        d_atr.shift(1).rename("prior_atr"),
        left_on="session_date",
        right_index=True,
        how="left",
    )
    f = f[
        (f["or_open"] > cfg.min_price)
        & (f["avg_vol_14"] >= cfg.min_avg_vol_14)
        & (f["prior_atr"] > cfg.min_atr_14)
    ]
    if f.empty:
        return []

    rows: list[OrScreenRow] = []
    for rec in f.itertuples(index=False):
        if rec.or_close > rec.or_open:
            side: Side | None = "long"
        elif rec.or_close < rec.or_open:
            side = "short"
        else:
            side = None
        rows.append(
            OrScreenRow(
                session_date=pd.Timestamp(rec.session_date),
                symbol=sym,
                rel_vol=float("nan"),
                or_open=float(rec.or_open),
                or_high=float(rec.or_high),
                or_low=float(rec.or_low),
                or_close=float(rec.or_close),
                or_volume=float(rec.or_volume),
                prior_atr=float(rec.prior_atr),
                avg_vol_14=float(rec.avg_vol_14),
                side=side,
            )
        )
    return rows


def build_or_screen(
    intra: Dict[str, pd.DataFrame],
    daily: Dict[str, pd.DataFrame],
    *,
    cfg: OrbZarattiniConfig,
    verbose: bool = False,
    apply_filters: bool = True,
) -> pd.DataFrame:
    """Per (session, symbol) opening-range features before top-N filter.

    If ``apply_filters`` is False, keep all sides with finite rel_vol (for sweep caches).
    """
    chunks: list[pd.DataFrame] = []
    n = len(intra)
    for k, sym in enumerate(sorted(intra.keys()), start=1):
        rows = _extract_or_rows(sym, intra[sym], daily.get(sym, pd.DataFrame()), cfg)
        if rows:
            chunks.append(pd.DataFrame([r.__dict__ for r in rows]))
        if verbose and k % 100 == 0:
            print(f"  screened {k}/{n} symbols …", flush=True)
    if not chunks:
        return pd.DataFrame()
    df = pd.concat(chunks, ignore_index=True)
    df["session_date"] = pd.to_datetime(df["session_date"])

    # Relative volume: OR volume / 14-day mean of prior OR volumes (paper Eq. 1).
    parts = []
    for sym, grp in df.groupby("symbol", sort=False):
        g = grp.sort_values("session_date").copy()
        g["rel_vol"] = g["or_volume"] / g["or_volume"].shift(1).rolling(cfg.or_lookback, min_periods=cfg.or_lookback).mean()
        parts.append(g)
    out = pd.concat(parts, ignore_index=True)
    out = out[np.isfinite(out["rel_vol"].astype(float))]
    out = out[out["side"].notna()]
    # OR geometry features for quality filters / sweeps.
    or_range = (out["or_high"] - out["or_low"]).astype(float)
    out["or_range"] = or_range
    out["or_range_atr"] = or_range / out["prior_atr"].astype(float).replace(0, np.nan)
    body = (out["or_close"] - out["or_open"]).abs().astype(float)
    out["body_frac"] = np.where(or_range > 0, body / or_range, np.nan)
    if apply_filters:
        if float(cfg.min_rel_vol) > 0:
            out = out[out["rel_vol"] >= float(cfg.min_rel_vol)]
        out = apply_or_quality_filters(out, cfg)
    return out.sort_values(["session_date", "rel_vol"], ascending=[True, False]).reset_index(drop=True)


def apply_or_quality_filters(screen: pd.DataFrame, cfg: OrbZarattiniConfig) -> pd.DataFrame:
    """Filter a prebuilt OR screen by RVOL / OR-range / body rules."""
    if screen.empty:
        return screen
    out = screen
    if float(cfg.min_rel_vol) > 0 and "rel_vol" in out.columns:
        out = out[out["rel_vol"] >= float(cfg.min_rel_vol)]
    if "or_range_atr" not in out.columns and {"or_high", "or_low", "prior_atr"}.issubset(out.columns):
        or_range = (out["or_high"] - out["or_low"]).astype(float)
        out = out.copy()
        out["or_range"] = or_range
        out["or_range_atr"] = or_range / out["prior_atr"].astype(float).replace(0, np.nan)
        body = (out["or_close"] - out["or_open"]).abs().astype(float)
        out["body_frac"] = np.where(or_range > 0, body / or_range, np.nan)
    if float(cfg.min_or_range_atr) > 0:
        out = out[out["or_range_atr"] >= float(cfg.min_or_range_atr)]
    if float(cfg.max_or_range_atr) > 0:
        out = out[out["or_range_atr"] <= float(cfg.max_or_range_atr)]
    if float(cfg.min_body_frac) > 0:
        out = out[out["body_frac"] >= float(cfg.min_body_frac)]
    return out.reset_index(drop=True)


def select_top_n_per_session(screen: pd.DataFrame, *, top_n: int) -> pd.DataFrame:
    if screen.empty:
        return screen
    return (
        screen.sort_values(["session_date", "rel_vol"], ascending=[True, False])
        .groupby("session_date", sort=True, group_keys=False)
        .head(int(top_n))
        .reset_index(drop=True)
    )


def _size_shares(
    equity: float,
    entry_px: float,
    stop_dist: float,
    *,
    cfg: OrbZarattiniConfig,
    n_positions: int,
) -> int:
    if not (np.isfinite(equity) and equity > 0 and np.isfinite(entry_px) and entry_px > 0 and stop_dist > 0):
        return 0
    risk_usd = equity * float(cfg.risk_per_trade)
    by_risk = int(risk_usd / stop_dist)
    # Portfolio gross cap: 4x equity split across concurrent day positions.
    lev_budget = equity * float(cfg.max_leverage) / max(1, n_positions)
    by_lev = int(lev_budget / entry_px)
    return max(0, min(by_risk, by_lev))


def _commission(shares: int, cfg: OrbZarattiniConfig) -> float:
    return 2.0 * abs(shares) * float(cfg.commission_per_share)


def simulate_orb_trade(
    intra: pd.DataFrame,
    pick: OrScreenRow | pd.Series,
    *,
    cfg: OrbZarattiniConfig,
    equity: float,
    n_positions: int,
) -> tuple[TradeRow | None, float]:
    """Return (trade, dollar_pnl) for one session/symbol."""
    sym = str(pick["symbol"] if isinstance(pick, pd.Series) else pick.symbol)
    session_date = pd.Timestamp(pick["session_date"] if isinstance(pick, pd.Series) else pick.session_date)
    side = pick["side"] if isinstance(pick, pd.Series) else pick.side
    or_high = float(pick["or_high"] if isinstance(pick, pd.Series) else pick.or_high)
    or_low = float(pick["or_low"] if isinstance(pick, pd.Series) else pick.or_low)
    prior_atr = float(pick["prior_atr"] if isinstance(pick, pd.Series) else pick.prior_atr)
    if side is None:
        return None, 0.0

    idx = pd.to_datetime(intra.index).tz_localize(None)
    ses = _session_key(idx)
    bar_n = _session_bar_index(idx)
    day_mask = ses == session_date
    if not day_mask.any():
        return None, 0.0

    o = intra["open"].astype(np.float64)
    h = intra["high"].astype(np.float64)
    l = intra["low"].astype(np.float64)
    c = intra["close"].astype(np.float64)

    or_bars = _or_bar_count(cfg)
    stop_dist = float(cfg.stop_atr_frac) * prior_atr
    slip = float(cfg.slippage_bps)
    entry_bar = None
    entry_px = None
    stop_px = None

    for i in np.flatnonzero(day_mask.values):
        # Entries only after the opening-range window completes.
        bn = int(bar_n.iloc[i])
        if bn < or_bars:
            continue
        if int(cfg.max_entry_bar) > 0 and bn > int(cfg.max_entry_bar):
            break
        hi = float(h.iloc[i])
        lo = float(l.iloc[i])
        prev_i = i - 1
        prev_in_day = prev_i >= 0 and day_mask.iloc[prev_i] and int(bar_n.iloc[prev_i]) >= or_bars
        prev_hi = float(h.iloc[prev_i]) if prev_in_day else float("nan")
        prev_lo = float(l.iloc[prev_i]) if prev_in_day else float("nan")
        if side == "long":
            if hi >= or_high and (not np.isfinite(prev_hi) or prev_hi < or_high):
                raw = max(or_high, float(o.iloc[i]))
                entry_px = _slip(raw, slip, "buy")
                stop_px = entry_px - stop_dist
                entry_bar = int(i)
                break
        else:
            if lo <= or_low and (not np.isfinite(prev_lo) or prev_lo > or_low):
                raw = min(or_low, float(o.iloc[i]))
                entry_px = _slip(raw, slip, "sell")
                stop_px = entry_px + stop_dist
                entry_bar = int(i)
                break

    if entry_bar is None or entry_px is None or stop_px is None:
        return None, 0.0

    shares = _size_shares(equity, entry_px, stop_dist, cfg=cfg, n_positions=n_positions)
    if shares <= 0:
        return None, 0.0

    exit_bar = None
    exit_px = None
    reason = "session_close"
    for j in range(entry_bar + 1, len(idx)):
        if not day_mask.iloc[j]:
            break
        hi = float(h.iloc[j])
        lo = float(l.iloc[j])
        cl = float(c.iloc[j])
        if side == "long" and lo <= stop_px:
            exit_px = _slip(stop_px, slip, "sell")
            reason = "stop_atr"
            exit_bar = j
            break
        if side == "short" and hi >= stop_px:
            exit_px = _slip(stop_px, slip, "buy")
            reason = "stop_atr"
            exit_bar = j
            break
        exit_bar = j
        exit_px = _slip(cl, slip, "sell" if side == "long" else "buy")

    if exit_bar is None:
        # Entry on last bar of session — flat at entry (no follow-through).
        exit_bar = entry_bar
        exit_px = _slip(float(c.iloc[entry_bar]), slip, "sell" if side == "long" else "buy")
        reason = "session_close"

    if exit_bar is None or exit_px is None:
        return None, 0.0

    if side == "long":
        gross = (exit_px - entry_px) * shares
        pnl_pct = exit_px / entry_px - 1.0
    else:
        gross = (entry_px - exit_px) * shares
        pnl_pct = entry_px / exit_px - 1.0 if exit_px > 0 else 0.0
    net = gross - _commission(shares, cfg)

    trade = TradeRow(
        variant=f"orb{int(cfg.or_minutes)}_{side}",
        symbol=sym,
        session_date=str(session_date.date()),
        entry_time=str(idx[entry_bar]),
        exit_time=str(idx[exit_bar]),
        entry_price=float(entry_px),
        exit_price=float(exit_px),
        pnl_pct=float(pnl_pct),
        bars_held=int(exit_bar - entry_bar),
        exit_reason=reason,
    )
    return trade, float(net)


def run_orb_backtest(
    intra: Dict[str, pd.DataFrame],
    daily: Dict[str, pd.DataFrame],
    *,
    cfg: OrbZarattiniConfig | None = None,
    return_start: pd.Timestamp | None = None,
    return_end: pd.Timestamp | None = None,
    verbose: bool = True,
    screen: pd.DataFrame | None = None,
) -> tuple[pd.Series, pd.DataFrame, dict]:
    cfg = cfg or zarattini_config()
    _ = _or_bar_count(cfg)  # validate divisibility early
    if screen is None:
        screen = build_or_screen(intra, daily, cfg=cfg, verbose=verbose, apply_filters=True)
        filtered = screen
    else:
        filtered = apply_or_quality_filters(screen, cfg)
        if float(cfg.min_rel_vol) > 0:
            filtered = filtered[filtered["rel_vol"] >= float(cfg.min_rel_vol)]
    picks = select_top_n_per_session(filtered, top_n=cfg.top_n)
    if picks.empty:
        return pd.Series(dtype=float), pd.DataFrame(), {}

    if return_start is not None:
        picks = picks[picks["session_date"] >= pd.Timestamp(return_start)]
    if return_end is not None:
        picks = picks[picks["session_date"] <= pd.Timestamp(return_end)]
    if picks.empty:
        return pd.Series(dtype=float), pd.DataFrame(), {}

    sessions = sorted(picks["session_date"].unique())
    equity = float(cfg.starting_capital)
    trades: list[TradeRow] = []
    daily_pnl: dict[pd.Timestamp, float] = {}

    for ses in sessions:
        day_picks = picks[picks["session_date"] == ses].copy()
        n_pos = len(day_picks)
        day_equity = float(cfg.starting_capital) if cfg.fixed_sizing else equity
        # Optional: allocate more risk to higher-RVOL names within the day.
        if cfg.rvol_rank_size and n_pos > 0 and "rel_vol" in day_picks.columns:
            rv = day_picks["rel_vol"].astype(float).clip(lower=1e-9)
            w = rv / rv.sum()
            day_picks = day_picks.assign(_risk_scale=(w * n_pos).clip(0.5, 2.0))
        else:
            day_picks = day_picks.assign(_risk_scale=1.0)
        day_net = 0.0
        for _, row in day_picks.iterrows():
            sym = str(row["symbol"])
            ib = intra.get(sym)
            if ib is None or ib.empty:
                continue
            sized_eq = day_equity * float(row["_risk_scale"])
            trade, net = simulate_orb_trade(ib, row, cfg=cfg, equity=sized_eq, n_positions=n_pos)
            if trade is None:
                continue
            trades.append(trade)
            day_net += net
        equity += day_net
        daily_pnl[pd.Timestamp(ses)] = day_net

    port = pd.Series(daily_pnl, dtype=np.float64).sort_index()
    # Convert dollar PnL to return on start-of-day equity for Sharpe helper.
    eq_curve = [float(cfg.starting_capital)]
    rets = []
    idx_ret = []
    for dt, pnl in port.items():
        base = eq_curve[-1]
        r = pnl / base if base > 0 else 0.0
        rets.append(r)
        idx_ret.append(dt)
        eq_curve.append(base + pnl)
    port_r = pd.Series(rets, index=pd.DatetimeIndex(idx_ret), dtype=np.float64)

    tdf = pd.DataFrame([t.__dict__ for t in trades])
    meta = {
        "starting_capital": float(cfg.starting_capital),
        "ending_equity": float(eq_curve[-1]),
        "total_return_pct": float((eq_curve[-1] / cfg.starting_capital - 1.0) * 100),
        "or_minutes": int(cfg.or_minutes),
        "bar_minutes": int(cfg.bar_minutes),
        "top_n": int(cfg.top_n),
        "min_rel_vol": float(cfg.min_rel_vol),
        "stop_atr_frac": float(cfg.stop_atr_frac),
        "min_or_range_atr": float(cfg.min_or_range_atr),
        "max_or_range_atr": float(cfg.max_or_range_atr),
        "min_body_frac": float(cfg.min_body_frac),
        "max_entry_bar": int(cfg.max_entry_bar),
        "rvol_rank_size": bool(cfg.rvol_rank_size),
        "n_sessions_traded": int(len(port)),
        "n_trades": int(len(trades)),
        "screen_rows": int(len(filtered)),
        "pick_rows": int(len(picks)),
        **portfolio_metrics(port_r),
        **trade_stats(tdf),
    }
    return port_r, tdf, meta


def beta_vs_spy(port_r: pd.Series, spy_daily_ret: pd.Series) -> float:
    pr = port_r.copy()
    pr.index = pd.to_datetime(pr.index).normalize()
    spy = spy_daily_ret.copy()
    spy.index = pd.to_datetime(spy.index).normalize()
    aligned = pd.concat([pr, spy], axis=1, join="inner").dropna()
    if len(aligned) < 10:
        return float("nan")
    y = aligned.iloc[:, 0].astype(float)
    x = aligned.iloc[:, 1].astype(float)
    xv = float(x.var(ddof=1))
    if xv < 1e-18:
        return float("nan")
    return float(np.cov(y, x, ddof=1)[0, 1] / xv)
