"""
Hourly swing-high / swing-low **breakout** research backtest (e.g. equities or listed ETFs).

Swings: Williams-style fractals (strict highs / lows with ``fractal_left`` / ``fractal_right``).
A swing is **confirmed** after the right side completes (no lookahead). **Significance**:
distance to the prior opposite swing must be at least ``min_atr_mult`` × **prior session's**
Wilder daily ATR (from daily OHLC aggregated from the hourly bars).

Entry: first bar after confirmation where price **crosses** the swing level (long: prior high
below level, this bar's high at or above; optionally require **close** through the level).
Fill at the swing level (stop price). Each swing level is traded **at most once**.

Exit: **time stop** — flat after ``hold_bars`` hourly bars in the trade (direction × bar
close-to-close returns).

Short sales ignore borrow/fees (research-style). Intended for Yahoo hourly data
(``DataLoader.fetch_intraday``; depth is vendor-limited, often ~730d for 60m).
"""

from __future__ import annotations

from dataclasses import dataclass
from typing import Literal

import numpy as np
import pandas as pd

SwingPortfolioBlend = Literal["equal", "inv_vol", "winners_equal"]


def _hourly_bar_returns(close: pd.Series) -> np.ndarray:
    """Close-to-close simple returns; avoids pandas 2.x ``pct_change`` pad deprecation."""
    c = close.astype(np.float64)
    try:
        r = c.pct_change(fill_method=None)
    except TypeError:
        r = c.pct_change()
    out = r.to_numpy(dtype=np.float64)
    out[0] = 0.0
    np.nan_to_num(out, copy=False, nan=0.0)
    return out


def _wilder_atr(high: pd.Series, low: pd.Series, close: pd.Series, period: int) -> pd.Series:
    p = int(period)
    prev = close.shift(1)
    tr = pd.concat(
        [(high - low).abs(), (high - prev).abs(), (low - prev).abs()],
        axis=1,
    ).max(axis=1)
    return tr.ewm(alpha=1.0 / float(p), min_periods=p, adjust=False).mean()


def _fractal_swing_high(high: np.ndarray, left: int, right: int) -> np.ndarray:
    n = len(high)
    L, R = int(left), int(right)
    out = np.zeros(n, dtype=bool)
    for i in range(L, n - R):
        v = high[i]
        ok = True
        for k in range(1, L + 1):
            if v <= high[i - k]:
                ok = False
                break
        if not ok:
            continue
        for k in range(1, R + 1):
            if v <= high[i + k]:
                ok = False
                break
        out[i] = ok
    return out


def _fractal_swing_low(low: np.ndarray, left: int, right: int) -> np.ndarray:
    n = len(low)
    L, R = int(left), int(right)
    out = np.zeros(n, dtype=bool)
    for i in range(L, n - R):
        v = low[i]
        ok = True
        for k in range(1, L + 1):
            if v >= low[i - k]:
                ok = False
                break
        if not ok:
            continue
        for k in range(1, R + 1):
            if v >= low[i + k]:
                ok = False
                break
        out[i] = ok
    return out


def _daily_atr_series_from_hourly(hourly: pd.DataFrame, atr_period: int) -> pd.Series:
    """One row per calendar session in index (date); ATR shifted by 1 day (causal for intraday)."""
    if "trade_date" not in hourly.columns:
        raise KeyError("hourly must have trade_date (use DataLoader.align_to_trading_days)")
    g = hourly.groupby("trade_date", sort=True)
    daily = pd.DataFrame(
        {
            "open": g["open"].first(),
            "high": g["high"].max(),
            "low": g["low"].min(),
            "close": g["close"].last(),
        }
    )
    atr = _wilder_atr(daily["high"], daily["low"], daily["close"], atr_period)
    # Use only information available before session D: ATR through prior session close.
    atr_causal = atr.shift(1)
    return atr_causal


def _merge_daily_atr_to_hourly(hourly: pd.DataFrame, daily_atr: pd.Series) -> np.ndarray:
    da = daily_atr.copy()
    da.index = pd.to_datetime(da.index).normalize()
    dates_norm = pd.to_datetime(hourly["trade_date"]).dt.normalize()
    return dates_norm.map(da).astype(np.float64).to_numpy()


def _session_extremes_cum(high: np.ndarray, low: np.ndarray, trade_dates: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
    """Per bar: running session high / low from the session open through this bar (causal)."""
    n = len(high)
    sh = np.empty(n, dtype=np.float64)
    sl = np.empty(n, dtype=np.float64)
    cur = None
    rh = -np.inf
    rl = np.inf
    for t in range(n):
        d = trade_dates[t]
        if cur is None or d != cur:
            cur = d
            rh = high[t]
            rl = low[t]
        else:
            rh = max(rh, high[t])
            rl = min(rl, low[t])
        sh[t] = rh
        sl[t] = rl
    return sh, sl


@dataclass
class SwingBreakoutHourlyConfig:
    fractal_left: int = 2
    fractal_right: int = 2
    daily_atr_period: int = 14
    min_atr_mult: float = 1.0
    hold_bars: int = 10
    long_side: bool = True
    short_side: bool = True
    #: If True, long requires ``close >= level`` and short requires ``close <= level`` on the entry bar.
    close_confirms_breakout: bool = False


def _trade_pnl_from_bars(ret: np.ndarray, entry_t: int, exit_bar: int, side: int) -> float:
    """Compound PnL over bars where position is open (matches shift(1)*ret accounting)."""
    lo = entry_t + 1
    hi = min(exit_bar + 2, len(ret))
    if lo >= hi:
        return 0.0
    chunk = ret[lo:hi]
    return float(np.prod(1.0 + float(side) * chunk) - 1.0)


def simulate_ticker_hourly(
    hourly: pd.DataFrame,
    cfg: SwingBreakoutHourlyConfig,
    *,
    record_trade_pnls: bool = False,
) -> tuple[pd.Series, int, list[float]]:
    """
    Returns (hourly simple return series, trade_count, trade_pnls).

    ``trade_pnls`` is a list of simple returns per **completed** round-trip (empty unless
    ``record_trade_pnls=True``). Open positions at end of sample are included as partial PnL.
    """
    if hourly.empty or len(hourly) < 20:
        return (
            pd.Series(0.0, index=hourly.index, dtype=np.float64),
            0,
            [],
        )

    df = hourly.sort_index().copy()
    if "trade_date" not in df.columns:
        raise KeyError("hourly must have trade_date")

    close = df["close"].to_numpy(dtype=np.float64)
    high = df["high"].to_numpy(dtype=np.float64)
    low = df["low"].to_numpy(dtype=np.float64)
    ret = _hourly_bar_returns(df["close"])

    L, R = int(cfg.fractal_left), int(cfg.fractal_right)
    hold = int(cfg.hold_bars)
    mult = float(cfg.min_atr_mult)

    daily_atr = _daily_atr_series_from_hourly(df, int(cfg.daily_atr_period))
    atr_h = _merge_daily_atr_to_hourly(df, daily_atr)

    sh = _fractal_swing_high(high, L, R)
    sl = _fractal_swing_low(low, L, R)

    n = len(df)
    micro = np.zeros(n, dtype=np.int8)
    in_pos = 0
    exit_bar = -1
    trades = 0
    trade_entry_t = -1
    trade_side = 0
    trade_pnls: list[float] = []

    # Untraded swing levels (FIFO). Tuple: (confirm_bar_index, level, kind 'H'/'L')
    pending_highs: list[tuple[int, float]] = []
    pending_lows: list[tuple[int, float]] = []

    last_sig_high = np.nan
    last_sig_low = np.nan
    td = pd.to_datetime(df["trade_date"]).dt.normalize().to_numpy()
    sess_h, sess_l = _session_extremes_cum(high, low, td)

    for t in range(n):
        if in_pos != 0 and t > exit_bar:
            if record_trade_pnls and trade_entry_t >= 0:
                trade_pnls.append(_trade_pnl_from_bars(ret, trade_entry_t, exit_bar, trade_side))
            trade_entry_t = -1
            in_pos = 0

        # Confirm fractals whose center i satisfies i + R == t
        if t >= L + R:
            i = t - R
            atr_i = atr_h[i] if i < len(atr_h) else np.nan
            if np.isfinite(atr_i) and atr_i > 0:
                if sh[i]:
                    lvl = float(high[i])
                    ok = False
                    if cfg.long_side:
                        if np.isfinite(last_sig_low):
                            ok = (lvl - last_sig_low) >= mult * atr_i
                        else:
                            ok = (lvl - sess_l[i]) >= mult * atr_i
                    if ok:
                        pending_highs.append((t, lvl))
                        last_sig_high = lvl
                if sl[i]:
                    lvl = float(low[i])
                    ok = False
                    if cfg.short_side:
                        if np.isfinite(last_sig_high):
                            ok = (last_sig_high - lvl) >= mult * atr_i
                        else:
                            ok = (sess_h[i] - lvl) >= mult * atr_i
                    if ok:
                        pending_lows.append((t, lvl))
                        last_sig_low = lvl

        if in_pos == 0:
            # Long breakouts (oldest pending high first)
            if cfg.long_side and t > 0:
                to_remove: list[int] = []
                for pi, (conf_bar, lvl) in enumerate(pending_highs):
                    if t <= conf_bar:
                        continue
                    cross = high[t] >= lvl and high[t - 1] < lvl
                    if cfg.close_confirms_breakout:
                        cross = cross and (close[t] >= lvl)
                    if cross:
                        in_pos = 1
                        exit_bar = t + hold - 1
                        trades += 1
                        trade_entry_t = t
                        trade_side = 1
                        to_remove.append(pi)
                        break
                if to_remove:
                    pending_highs = [p for j, p in enumerate(pending_highs) if j not in to_remove]

            if in_pos == 0 and cfg.short_side and t > 0:
                to_remove2: list[int] = []
                for pi, (conf_bar, lvl) in enumerate(pending_lows):
                    if t <= conf_bar:
                        continue
                    cross = low[t] <= lvl and low[t - 1] > lvl
                    if cfg.close_confirms_breakout:
                        cross = cross and (close[t] <= lvl)
                    if cross:
                        in_pos = -1
                        exit_bar = t + hold - 1
                        trades += 1
                        trade_entry_t = t
                        trade_side = -1
                        to_remove2.append(pi)
                        break
                if to_remove2:
                    pending_lows = [p for j, p in enumerate(pending_lows) if j not in to_remove2]

        micro[t] = in_pos

    if record_trade_pnls and in_pos != 0 and trade_entry_t >= 0:
        lo = trade_entry_t + 1
        if lo < n:
            chunk = ret[lo:n]
            trade_pnls.append(float(np.prod(1.0 + float(trade_side) * chunk) - 1.0))

    micro_s = pd.Series(micro, index=df.index, dtype=np.float64)
    strat_ret = micro_s.shift(1).fillna(0.0) * pd.Series(ret, index=df.index)
    return strat_ret.astype(np.float64), trades, trade_pnls


def generate_portfolio_returns(
    hourly_by_ticker: dict[str, pd.DataFrame],
    cfg: SwingBreakoutHourlyConfig,
    *,
    blend: SwingPortfolioBlend = "equal",
    vol_window: int = 20,
    winner_epsilon: float = 0.0,
    verbose: bool = True,
) -> tuple[pd.Series, dict[str, int]]:
    """
    Combine per-ticker strategy returns on a common timeline.

    * ``equal`` — simple mean of aligned series (missing bar → 0 before mean).
    * ``inv_vol`` — each bar, weight names by ``1 / trailing hourly vol`` of **underlying** closes
      (``vol_window``, shifted one bar; row renormalized to sum to 1).
    * ``winners_equal`` — **in-sample**: keep only tickers whose full-sample compounded strategy
      return is ``> winner_epsilon``; equal-weight mean among those (others dropped). Biased if
      used for the same period that selected winners; use for diagnostics or a train/hold split.
    """
    if not hourly_by_ticker:
        raise ValueError("hourly_by_ticker is empty")

    b = str(blend).lower().strip()
    if b not in ("equal", "inv_vol", "winners_equal"):
        raise ValueError(f"blend must be equal, inv_vol, or winners_equal; got {blend!r}")

    tickers = list(hourly_by_ticker.keys())
    series_list: list[pd.Series] = []
    counts: dict[str, int] = {}
    for tkr in tickers:
        h = hourly_by_ticker[tkr]
        r, ntr, _ = simulate_ticker_hourly(h, cfg)
        series_list.append(r.rename(tkr))
        counts[tkr] = ntr

    master_idx = series_list[0].index
    for s in series_list[1:]:
        master_idx = master_idx.union(s.index).sort_values()

    aligned = [s.reindex(master_idx).fillna(0.0).astype(np.float64) for s in series_list]
    r_mat = pd.concat(aligned, axis=1)
    r_mat.columns = tickers
    r_np = r_mat.to_numpy(dtype=np.float64)
    t_n, n_sym = r_np.shape

    if b == "equal":
        port_arr = np.mean(r_np, axis=1)
    elif b == "winners_equal":
        growth = np.prod(1.0 + r_np, axis=0)
        cums = growth - 1.0
        mask = cums > float(winner_epsilon)
        if not np.any(mask):
            if verbose:
                print(
                    "  [winners_equal] No tickers above epsilon; falling back to full equal-weight."
                )
            port_arr = np.mean(r_np, axis=1)
        else:
            sub = r_np[:, mask]
            port_arr = np.mean(sub, axis=1)
            kept = [tickers[i] for i in range(n_sym) if mask[i]]
            dropped = [tickers[i] for i in range(n_sym) if not mask[i]]
            if verbose:
                print(
                    f"  [winners_equal] Kept {len(kept)}/{n_sym} (cum > {winner_epsilon:g}): "
                    f"{kept}  |  dropped: {dropped}"
                )
    else:
        vw = max(2, int(vol_window))
        inv = np.zeros((t_n, n_sym), dtype=np.float64)
        eps = 1e-8
        for j, tkr in enumerate(tickers):
            close = hourly_by_ticker[tkr]["close"].astype(np.float64).reindex(master_idx).ffill()
            ar = pd.Series(_hourly_bar_returns(close), index=master_idx, dtype=np.float64)
            vol = ar.rolling(window=vw, min_periods=vw).std(ddof=1).shift(1)
            v = vol.to_numpy(dtype=np.float64)
            inv[:, j] = np.where(np.isfinite(v) & (v > eps), 1.0 / v, 0.0)
        row_sum = inv.sum(axis=1, keepdims=True)
        w = np.divide(inv, row_sum, out=np.zeros_like(inv), where=row_sum > eps)
        dead = (row_sum.ravel() < eps)
        if np.any(dead):
            w[dead, :] = 1.0 / float(n_sym)
        port_arr = (r_np * w).sum(axis=1)

    port = pd.Series(port_arr, index=master_idx, dtype=np.float64)
    port.name = "swing_breakout_hourly_ret"

    if verbose:
        n = int(sum(counts.values()))
        print(
            f"  Hourly swing breakout: {len(hourly_by_ticker)} names | blend={b} | "
            f"total trades (all tickers) ≈ {n} | hold={cfg.hold_bars} bars | "
            f"fractal L/R={cfg.fractal_left}/{cfg.fractal_right} | min move ≥ {cfg.min_atr_mult}× daily ATR"
        )
        if b == "inv_vol":
            print(f"  inv_vol: vol_window={vw} (underlying hourly close returns, causal)")
    return port, counts


def generate_equal_weight_portfolio(
    hourly_by_ticker: dict[str, pd.DataFrame],
    cfg: SwingBreakoutHourlyConfig,
    *,
    verbose: bool = True,
) -> tuple[pd.Series, dict[str, int]]:
    """Backward-compatible alias for :func:`generate_portfolio_returns` with ``blend=equal``."""
    return generate_portfolio_returns(hourly_by_ticker, cfg, blend="equal", verbose=verbose)
