#!/usr/bin/env python3
"""
Prototype **market-neutral pairs stat-arb** book (daily bars).

Scans S&P 100 for cointegrated pairs (Engle–Granger + Holm–Bonferroni + Hurst/half-life),
runs causal OLS hedge + spread z-score mean reversion on each pair, equal-weight portfolio.

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 \\
      .venv/bin/python RenTech/strategy_stack/run_pairs_statarb_standard.py \\
      --start 2016-01-04 --end 2026-04-02 --capital 100000 --top-k 5
"""

from __future__ import annotations

import argparse
import json
import sys
from pathlib import Path

import numpy as np
import pandas as pd

_REPO = Path(__file__).resolve().parents[2]
if str(_REPO) not in sys.path:
    sys.path.insert(0, str(_REPO))

import universe_scanner as us  # type: ignore[import-not-found]

from RenTech.strategy_stack.data_loader import DataLoader
from RenTech.strategy_stack.main import (
    _compute_daily_backtest_features,
    find_cointegrated_pairs_with_fallback,
    vectorized_strategy_returns,
)
from RenTech.strategy_stack.statarb_engine import PairStatArbEngine

LOGS = _REPO / "RenTech" / "data" / "logs"
DEFAULT_OUT_PREFIX = LOGS / "pairs_statarb_standard"


def _close_frame(df: pd.DataFrame) -> pd.DataFrame:
    out = df.copy()
    if "close" not in out.columns:
        if "Close" in out.columns:
            out["close"] = out["Close"].astype(np.float64)
        else:
            raise KeyError("DataFrame missing close/Close column")
    return out[["close"]]


def _backtest_pair_daily(
    df_y: pd.DataFrame,
    df_x: pd.DataFrame,
    *,
    window: int,
    entry_z: float,
    exit_z: float,
    hedge_window: int,
    min_hedge_obs: int,
) -> pd.Series:
    eng = PairStatArbEngine(
        window=int(window),
        entry_z=float(entry_z),
        exit_z=float(exit_z),
        hedge_window=int(hedge_window) if hedge_window > 0 else None,
        min_hedge_obs=int(min_hedge_obs),
    )
    out = eng.transform(_close_frame(df_y), _close_frame(df_x))
    return vectorized_strategy_returns(out, "micro_position", ret_col="basket_ret").fillna(0.0)


def _beta_vs_spy(r: pd.Series, spy_r: pd.Series) -> tuple[float, float]:
    aligned = pd.DataFrame({"s": r, "spy": spy_r.reindex(r.index).fillna(0.0)}).dropna()
    if len(aligned) < 10:
        return float("nan"), float("nan")
    rho = float(aligned.corr().iloc[0, 1])
    sv = float(aligned["spy"].var())
    beta = float(aligned["s"].cov(aligned["spy"]) / sv) if sv > 1e-14 else float("nan")
    return beta, rho


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
    ap.add_argument("--start", default="2016-01-04")
    ap.add_argument("--end", default="2026-04-02")
    ap.add_argument("--lookback-years", type=float, default=10.0, help="Universe download window")
    ap.add_argument("--top-k", type=int, default=5, help="Max non-overlapping pairs in book")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT_PREFIX)
    ap.add_argument("--z-window", type=int, default=40)
    ap.add_argument("--z-entry", type=float, default=1.5)
    ap.add_argument("--z-exit", type=float, default=0.0)
    ap.add_argument("--hedge-window", type=int, default=60, help="Rolling OLS window (0=expanding)")
    ap.add_argument("--min-hedge-obs", type=int, default=40)
    ap.add_argument(
        "--pair-y",
        default="",
        help="Optional fixed leg Y (skip scan if both --pair-y and --pair-x set)",
    )
    ap.add_argument("--pair-x", default="", help="Optional fixed hedge leg X")
    ap.add_argument("--refresh-cache", action="store_true")
    args = ap.parse_args()

    cap = float(args.capital)
    t0 = pd.Timestamp(args.start)
    t1 = pd.Timestamp(args.end) if args.end.strip() else None

    print(f"Loading S&P 100 daily panel (lookback {args.lookback_years}y) …", flush=True)
    tickers = us.get_sp100_tickers(universe="sp100")
    raw = us.download_and_cache_data(
        tickers,
        timeframe="1d",
        lookback_years=float(args.lookback_years),
        max_age_hours=0.0 if args.refresh_cache else 24.0,
    )

    pair_y_in = str(args.pair_y).strip().upper()
    pair_x_in = str(args.pair_x).strip().upper()
    if bool(pair_y_in) ^ bool(pair_x_in):
        raise SystemExit("Pass both --pair-y and --pair-x, or neither.")

    coint_label: str | None = None
    if pair_y_in and pair_x_in:
        pairs_df = pd.DataFrame(
            [{"Ticker_A": pair_y_in, "Ticker_B": pair_x_in, "Adjusted_P_Value": 0.0}]
        )
        print(f"Fixed pair: {pair_y_in} (Y) vs {pair_x_in} (X)", flush=True)
    else:
        print("Scanning for cointegrated pairs …", flush=True)
        pairs_df, coint_label = find_cointegrated_pairs_with_fallback(us, raw)
        if pairs_df.empty:
            raise SystemExit(
                "No tradable cointegrated pairs found. Try --pair-y QQQ --pair-x SPY "
                "or relax scan via universe_scanner thresholds."
            )
        if coint_label:
            print(f"  Scan tier: {coint_label}", flush=True)
        pairs_df = pairs_df.head(int(args.top_k))
        print(f"  Selected {len(pairs_df)} pair(s):", flush=True)
        for _, row in pairs_df.iterrows():
            print(
                f"    {row['Ticker_A']} / {row['Ticker_B']}  "
                f"p_adj={float(row.get('Adjusted_P_Value', float('nan'))):.4f}",
                flush=True,
            )

    pair_rets: dict[str, pd.Series] = {}
    pair_meta: list[dict] = []
    for _, row in pairs_df.iterrows():
        y_t = str(row["Ticker_A"])
        x_t = str(row["Ticker_B"])
        if y_t not in raw or x_t not in raw:
            print(f"  Skip {y_t}/{x_t}: missing cache data", flush=True)
            continue
        try:
            r = _backtest_pair_daily(
                raw[y_t],
                raw[x_t],
                window=int(args.z_window),
                entry_z=float(args.z_entry),
                exit_z=float(args.z_exit),
                hedge_window=int(args.hedge_window),
                min_hedge_obs=int(args.min_hedge_obs),
            )
            r.index = pd.to_datetime(r.index).tz_localize(None)
            pair_rets[f"{y_t}_{x_t}"] = r
            pair_meta.append({"ticker_y": y_t, "ticker_x": x_t, "n_days": int(len(r))})
        except Exception as exc:
            print(f"  Skip {y_t}/{x_t}: {exc}", flush=True)

    if not pair_rets:
        raise SystemExit("No pair backtests succeeded.")

    panel = pd.DataFrame(pair_rets).fillna(0.0)
    port_r = panel.mean(axis=1)
    mask = port_r.index >= t0
    if t1 is not None:
        mask &= port_r.index <= t1
    r = port_r.loc[mask].fillna(0.0).astype(np.float64)

    pnl = r * cap
    eq_usd = cap * (1.0 + r).cumprod()
    n = len(r)
    years = n / 252.0
    end_eq = float(eq_usd.iloc[-1])
    tot = end_eq / cap - 1.0
    cagr = (end_eq / cap) ** (1.0 / years) - 1.0 if years > 0 else float("nan")
    max_dd = float((eq_usd / eq_usd.cummax() - 1.0).min())
    sd = float(r.std(ddof=1)) if n > 1 else float("nan")
    sharpe = float(r.mean() / sd * np.sqrt(252.0)) if sd > 1e-12 else float("nan")

    spy_df = _compute_daily_backtest_features(
        DataLoader().fetch_daily("SPY", period="max")
    )
    spy_r = spy_df["ret"].astype(float)
    spy_r.index = pd.to_datetime(spy_r.index).tz_localize(None)
    beta, rho = _beta_vs_spy(r, spy_r)

    # Rough margin: 2 legs × ~50% notional per pair × n_pairs equal weight
    n_pairs = len(pair_rets)
    margin_usd = eq_usd * 0.25 * min(2.0, 0.5 * n_pairs)

    prefix = args.out_prefix.expanduser().resolve()
    prefix.parent.mkdir(parents=True, exist_ok=True)
    daily_path = Path(f"{prefix}_daily.csv")
    pairs_path = Path(f"{prefix}_pairs.csv")
    meta_path = Path(f"{prefix}_meta.json")

    pd.DataFrame(
        {
            "date": r.index.strftime("%Y-%m-%d"),
            "daily_ret": r.values,
            "daily_pnl_usd": pnl.values,
            "equity_usd": eq_usd.values,
            "margin_usd": margin_usd.values,
        }
    ).to_csv(daily_path, index=False)
    pairs_df.to_csv(pairs_path, index=False)

    cmd = (
        f"cd {_REPO} && PYTHONUNBUFFERED=1 .venv/bin/python "
        f"RenTech/strategy_stack/run_pairs_statarb_standard.py "
        f"--start {args.start} --end {args.end} --top-k {args.top_k} --capital {cap:.0f}"
    )
    meta = {
        "strategy": "pairs_statarb",
        "token": "pairs_statarb",
        "market_neutral": True,
        "n_pairs": len(pair_rets),
        "pairs": pair_meta,
        "coint_scan_tier": coint_label,
        "z_window": int(args.z_window),
        "z_entry": float(args.z_entry),
        "z_exit": float(args.z_exit),
        "hedge_window": int(args.hedge_window),
        "capital_usd": cap,
        "start": str(r.index.min().date()),
        "end": str(r.index.max().date()),
        "n_trading_days": int(n),
        "total_return_pct": round(tot * 100.0, 2),
        "cagr_pct": round(cagr * 100.0, 2),
        "sharpe_daily": round(sharpe, 3),
        "max_drawdown_pct": round(max_dd * 100.0, 2),
        "beta_vs_spy": round(beta, 3),
        "corr_vs_spy": round(rho, 3),
        "command": cmd,
        "daily_csv": str(daily_path),
        "pairs_csv": str(pairs_path),
    }
    meta_path.write_text(json.dumps(meta, indent=2) + "\n", encoding="utf-8")

    print(
        f"\nPairs stat-arb ({len(pair_rets)} pairs): return {meta['total_return_pct']:+.1f}%  "
        f"Sharpe {sharpe:.2f}  maxDD {meta['max_drawdown_pct']:.1f}%  "
        f"β(SPY) {beta:.2f}  ρ(SPY) {rho:.2f}",
        flush=True,
    )
    print(f"Wrote {daily_path}", flush=True)
    print(f"Wrote {pairs_path}", flush=True)


if __name__ == "__main__":
    main()
