"""
Parameter search for momentum-filtered statistical arbitrage (vectorized stack).

Train / test split is **time-ordered** (first ``train_frac`` / remainder of intraday
bars). For each evaluation, **daily** history is clipped to the last timestamp in the
intraday slice so the regime filter does not use closes **after** the segment being
scored (avoid lookahead in the macro series).

Primary objective series: ``filtered_ret`` (momentum-gated positions × bar return).

Run from repository root::

    .venv/bin/python RenTech/strategy_stack/optimizer.py
    .venv/bin/python RenTech/strategy_stack/optimizer.py --hedge-ticker SPY --heatmap-path RenTech/strategy_stack/opt_heatmap.png
"""

from __future__ import annotations

import argparse
import itertools
import os
import sys
from dataclasses import dataclass
from typing import Any, Mapping

import numpy as np
import pandas as pd
from tqdm import tqdm

# Repository root (parent of RenTech/)
_REPO_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", ".."))
if _REPO_ROOT not in sys.path:
    sys.path.insert(0, _REPO_ROOT)

from RenTech.strategy_stack.data_loader import DataLoader
from RenTech.strategy_stack.main import vectorized_strategy_returns
from RenTech.strategy_stack.momentum_filter import MomentumFilter
from RenTech.strategy_stack.statarb_engine import PairStatArbEngine, StatArbEngine
from RenTech.strategy_stack.strategy_orchestrator import StrategyOrchestrator


@dataclass
class DataBundle:
    """Aligned daily regime series and intraday legs for optimization / backtests."""

    daily: pd.DataFrame
    intra_y: pd.DataFrame
    intra_x: pd.DataFrame | None


def clip_daily_to_intraday_end(daily: pd.DataFrame, intra_last_ts: pd.Timestamp) -> pd.DataFrame:
    """
    Keep only daily rows up to the calendar date of ``intra_last_ts`` (inclusive).

    Ensures regime features for intraday bars in a segment do not depend on future
    daily closes beyond that segment.
    """
    if daily.empty:
        return daily
    di = pd.to_datetime(daily.index, utc=True).tz_localize(None).normalize()
    end = pd.to_datetime(intra_last_ts, utc=True).tz_localize(None).normalize()
    return daily.loc[di <= end]


def min_hedge_obs_for_window(z_window: int) -> int:
    """Minimum training observations for causal OLS before spread / z-score (target: z_window + 10)."""
    return int(z_window) + 10


def run_filtered_backtest(
    *,
    intra_y: pd.DataFrame,
    intra_x: pd.DataFrame | None,
    daily: pd.DataFrame,
    statarb_window: int,
    z_entry: float,
    z_exit: float,
    mom_window: int,
    aqr_lookback: int = 252,
    aqr_skip: int = 21,
    regime_shift_sessions: int = 0,
) -> pd.DataFrame:
    """
    Vectorized pipeline: stat-arb signal → ``StrategyOrchestrator.merge_and_gate`` →
    ``filtered_ret``. Uses the same components as ``main.run_pipeline`` without fetching data.
    """
    if intra_y.empty:
        raise ValueError("intra_y is empty")
    if "close" not in daily.columns:
        raise KeyError("daily must have 'close'")
    if "trade_date" not in intra_y.columns:
        raise KeyError("intra_y must have 'trade_date' (use DataLoader.align_to_trading_days)")

    daily_clip = clip_daily_to_intraday_end(daily, intra_y.index.max())
    if daily_clip.empty:
        raise ValueError("daily_clip empty after end-date clip; check daily vs intraday ranges")

    # AQR 12-minus-1 momentum gate.
    mf = MomentumFilter(sma_window=mom_window, aqr_lookback=aqr_lookback, aqr_skip=aqr_skip)
    daily_regime = mf.transform(daily_clip["close"])

    sw = int(statarb_window)
    min_h = min_hedge_obs_for_window(sw)

    if intra_x is not None:
        peng = PairStatArbEngine(
            window=sw,
            entry_z=z_entry,
            exit_z=z_exit,
            hedge_window=sw,
            min_hedge_obs=min_h,
        )
        intra = peng.transform(intra_y, intra_x)
        intra["ret"] = intra["basket_ret"]
    else:
        eng = StatArbEngine(window=sw, entry_z=z_entry, exit_z=z_exit)
        intra = eng.transform(intra_y)
        intra["ret"] = intra["close"].pct_change()

    orch = StrategyOrchestrator(regime_shift_sessions=regime_shift_sessions)
    intra = orch.merge_and_gate(intra, daily_regime, trade_date_col="trade_date")
    intra["filtered_ret"] = vectorized_strategy_returns(intra, "orchestrated_position")
    return intra


def compute_metrics(
    returns: pd.Series,
    *,
    bars_per_year_sqrt: float | None = None,
) -> dict[str, float]:
    """
    Cumulative return, approximate Sharpe (same scaling as ``main.summarize`` for hourly),
    and max drawdown on compounded equity.
    """
    if bars_per_year_sqrt is None:
        bars_per_year_sqrt = float(np.sqrt(252 * 6))
    r = returns.fillna(0.0).astype(np.float64)
    if len(r) < 2 or r.std(ddof=1) < 1e-12:
        return {"cum_return": 0.0, "sharpe": float("nan"), "max_drawdown": 0.0}
    mu = float(r.mean())
    sd = float(r.std(ddof=1))
    sharpe = mu / sd * bars_per_year_sqrt if sd > 0 else float("nan")
    cum = float((1.0 + r).prod() - 1.0)
    eq = (1.0 + r).cumprod()
    peak = eq.cummax()
    mdd = float(((eq - peak) / peak).min())
    return {"cum_return": cum, "sharpe": sharpe, "max_drawdown": mdd}


def evaluate_params(
    params: Mapping[str, Any],
    data: Mapping[str, Any],
    *,
    aqr_lookback: int = 252,
    aqr_skip: int = 21,
    regime_shift_sessions: int = 0,
) -> dict[str, float]:
    """
    Run the vectorized momentum-filtered backtest for one parameter dict.

    Parameters
    ----------
    params
        Keys: ``statarb_window``, ``z_entry``, ``z_exit``, ``mom_window``.
    data
        Keys: ``daily`` (OHLCV with ``close``), ``intra_y``, optional ``intra_x`` for pairs.

    Returns
    -------
    dict with ``cum_return``, ``sharpe``, ``max_drawdown`` on ``filtered_ret``.
    """
    intra = run_filtered_backtest(
        intra_y=data["intra_y"],
        intra_x=data.get("intra_x"),
        daily=data["daily"],
        statarb_window=int(params["statarb_window"]),
        z_entry=float(params["z_entry"]),
        z_exit=float(params["z_exit"]),
        mom_window=int(params["mom_window"]),
        aqr_lookback=aqr_lookback,
        aqr_skip=aqr_skip,
        regime_shift_sessions=regime_shift_sessions,
    )
    return compute_metrics(intra["filtered_ret"])


def time_ordered_train_test_split(
    df: pd.DataFrame,
    train_frac: float = 0.7,
) -> tuple[pd.DataFrame, pd.DataFrame]:
    """Sort by index; first ``train_frac`` rows = train, remainder = test."""
    if not 0.0 < train_frac < 1.0:
        raise ValueError("train_frac must be in (0, 1)")
    y = df.sort_index()
    n = len(y)
    if n < 10:
        raise ValueError(f"Need more rows for split; got {n}")
    cut = max(1, min(n - 1, int(n * train_frac)))
    return y.iloc[:cut].copy(), y.iloc[cut:].copy()


def align_intra_x_to_y(intra_y: pd.DataFrame, intra_x: pd.DataFrame) -> pd.DataFrame:
    """Align hedge leg to ``intra_y`` index (``PairStatArbEngine`` inner-joins on overlap)."""
    return intra_x.reindex(intra_y.index)


def load_intraday_bundle(
    regime_ticker: str = "SPY",
    trade_ticker: str = "QQQ",
    hedge_ticker: str | None = None,
    daily_period: str = "5y",
    intraday_period: str = "730d",
    intraday_interval: str = "1h",
) -> DataBundle:
    """Fetch and align data (same conventions as ``main.run_pipeline``)."""
    loader = DataLoader()
    daily = loader.fetch_daily(regime_ticker, period=daily_period)
    if daily.empty:
        raise RuntimeError(f"No daily data for {regime_ticker}")

    iv = "60m" if intraday_interval in ("1h", "60m") else intraday_interval
    intra_y = loader.fetch_intraday(trade_ticker, interval=iv, period=intraday_period)  # type: ignore[arg-type]
    if intra_y.empty:
        raise RuntimeError(f"No intraday data for {trade_ticker}")
    intra_y = loader.align_to_trading_days(intra_y)

    intra_x: pd.DataFrame | None = None
    if hedge_ticker:
        intra_x = loader.fetch_intraday(hedge_ticker, interval=iv, period=intraday_period)  # type: ignore[arg-type]
        if intra_x.empty:
            raise RuntimeError(f"No intraday data for hedge {hedge_ticker}")
        intra_x = loader.align_to_trading_days(intra_x)
        intra_x = align_intra_x_to_y(intra_y, intra_x)

    return DataBundle(daily=daily, intra_y=intra_y, intra_x=intra_x)


def data_dict_for_segment(
    bundle: DataBundle,
    intra_y_seg: pd.DataFrame,
) -> dict[str, Any]:
    """Build evaluate_params ``data`` mapping; align hedge leg to segment index."""
    d: dict[str, Any] = {"daily": bundle.daily, "intra_y": intra_y_seg}
    if bundle.intra_x is not None:
        d["intra_x"] = align_intra_x_to_y(intra_y_seg, bundle.intra_x)
    return d


def plot_is_sharpe_heatmap(
    is_df: pd.DataFrame,
    *,
    save_path: str | None = None,
    title: str = "In-sample Sharpe (mean over z_exit × mom_window)",
) -> None:
    """
    Heatmap: rows = statarb_window, cols = z_entry; value = mean is_sharpe over
    other grid dimensions (smoother than a single spike).
    """
    try:
        import matplotlib.pyplot as plt
    except ImportError as e:
        raise ImportError("matplotlib required for heatmap. pip install matplotlib") from e

    sub = is_df.groupby(["statarb_window", "z_entry"], as_index=False)["is_sharpe"].mean()
    pivot = sub.pivot(index="statarb_window", columns="z_entry", values="is_sharpe")

    fig, ax = plt.subplots(figsize=(8, 4))
    try:
        import seaborn as sns  # type: ignore[import-untyped]

        sns.heatmap(pivot, annot=True, fmt=".2f", cmap="RdYlGn", ax=ax, center=0.0)
    except ImportError:
        im = ax.imshow(pivot.to_numpy(dtype=float), aspect="auto", cmap="RdYlGn", vmin=-1, vmax=1)
        ax.set_xticks(np.arange(pivot.shape[1]) + 0.5, labels=pivot.columns)
        ax.set_yticks(np.arange(pivot.shape[0]) + 0.5, labels=pivot.index)
        for i in range(pivot.shape[0]):
            for j in range(pivot.shape[1]):
                v = pivot.iloc[i, j]
                ax.text(j, i, f"{v:.2f}", ha="center", va="center", fontsize=8)
        fig.colorbar(im, ax=ax)

    ax.set_title(title)
    ax.set_xlabel("z_entry")
    ax.set_ylabel("statarb_window")
    fig.tight_layout()
    if save_path:
        fig.savefig(save_path, dpi=150)
        print(f"Saved heatmap → {save_path}")
    else:
        plt.show()
    plt.close(fig)


def run_optimization(
    bundle: DataBundle,
    train_frac: float = 0.7,
    top_k: int = 3,
    heatmap_path: str | None = None,
    plot_heatmap: bool = True,
    aqr_lookback: int = 252,
    aqr_skip: int = 21,
    regime_shift_sessions: int = 0,
) -> tuple[pd.DataFrame, None]:
    """
    Grid-search in-sample Sharpe on the training slice of ``bundle.intra_y``.

    Returns
    -------
    is_results
        All parameter combinations with ``is_sharpe`` (and other ``is_*`` metrics),
        sorted by ``is_sharpe`` descending.
    None
        Placeholder for a future comparison table; pipeline mode only uses ``is_results``.
    """
    _ = top_k  # reserved for future top-k OOS workflow
    train_y, _test_y = time_ordered_train_test_split(bundle.intra_y, train_frac=train_frac)
    data_train = data_dict_for_segment(bundle, train_y)

    z_window_grid = [20, 50, 100]
    z_entry_grid = [1.5, 2.0, 2.5]
    z_exit_grid = [0.0, 0.5]
    mom_window_grid = [100, 200]

    combos = list(
        itertools.product(z_window_grid, z_entry_grid, z_exit_grid, mom_window_grid)
    )
    rows: list[dict[str, Any]] = []
    for statarb_window, z_entry, z_exit, mom_window in tqdm(combos, desc="grid search"):
        p = {
            "statarb_window": statarb_window,
            "z_entry": z_entry,
            "z_exit": z_exit,
            "mom_window": mom_window,
        }
        m = evaluate_params(
            p,
            data_train,
            aqr_lookback=aqr_lookback,
            aqr_skip=aqr_skip,
            regime_shift_sessions=regime_shift_sessions,
        )
        rows.append({**p, **{f"is_{k}": v for k, v in m.items()}})

    is_res = pd.DataFrame(rows)
    is_res = is_res.sort_values("is_sharpe", ascending=False, na_position="last")
    if plot_heatmap:
        plot_is_sharpe_heatmap(is_res, save_path=heatmap_path)
    return is_res, None


def main() -> None:
    p = argparse.ArgumentParser(description=__doc__)
    p.add_argument("--regime-ticker", default="SPY")
    p.add_argument("--trade-ticker", default="QQQ")
    p.add_argument("--hedge-ticker", default="", help="If set, pair mode (needs statsmodels)")
    p.add_argument("--daily-period", default="5y")
    p.add_argument("--intraday-period", default="730d")
    p.add_argument("--interval", default="1h", choices=("1h", "60m", "15m", "30m"))
    p.add_argument("--train-frac", type=float, default=0.7)
    p.add_argument("--top-k", type=int, default=3)
    p.add_argument("--heatmap-path", default="", help="PNG path (empty = plt.show())")
    p.add_argument("--no-heatmap", action="store_true", help="Skip heatmap figure")
    p.add_argument("--aqr-lookback", type=int, default=252)
    p.add_argument("--aqr-skip", type=int, default=21)
    args = p.parse_args()

    os.chdir(_REPO_ROOT)
    hedge = args.hedge_ticker.strip() or None

    print("Loading data …")
    bundle = load_intraday_bundle(
        regime_ticker=args.regime_ticker,
        trade_ticker=args.trade_ticker,
        hedge_ticker=hedge,
        daily_period=args.daily_period,
        intraday_period=args.intraday_period,
        intraday_interval=args.interval,
    )
    print(f"Intraday bars: {len(bundle.intra_y):,} | pair={hedge is not None}")

    print(f"Grid search in-sample (train_frac={args.train_frac}) …")
    is_res, comp = run_optimization(
        bundle,
        train_frac=args.train_frac,
        top_k=args.top_k,
        heatmap_path=args.heatmap_path or None,
        plot_heatmap=not args.no_heatmap,
        aqr_lookback=args.aqr_lookback,
        aqr_skip=args.aqr_skip,
    )

    pd.set_option("display.max_columns", 20)
    pd.set_option("display.width", 120)
    print("\n=== In-sample grid (top Sharpe) ===")
    print(is_res.head(10).to_string(index=False))
    if comp is not None:
        print("\n=== Top parameter sets: Train vs Test ===")
        print(comp.to_string(index=False))
    print(f"\n(In-sample grid rows: {len(is_res)})")


if __name__ == "__main__":
    main()
