"""Shared setup + metrics for equity momentum backtests (S&P 500, Russell 3000)."""

from __future__ import annotations

import json
from pathlib import Path
from typing import Any

import numpy as np
import pandas as pd

from RenTech.strategy_stack.data_loader import DataLoader
from RenTech.strategy_stack.main import (
    _compute_daily_backtest_features,
    _parse_daily_period_years,
    _wiki_sector_map_for_tickers,
)
from RenTech.strategy_stack.ml_momentum_engine import fetch_sp500_constituents_with_sectors
from RenTech.strategy_stack.equity_universe_loaders import fetch_russell3000_tickers
from RenTech.strategy_stack.run_sp500_dip_standard import _filter_equity_by_history
from RenTech.strategy_stack.sp500_momentum_index import Sp500MomentumConfig, Sp500MomentumIndex
from RenTech.strategy_stack.sp500_pit_universe import (
    build_market_cap_panel,
    load_shares_outstanding_dict,
    sp500_membership_mask,
    sp500_union_tickers,
)


def load_equity_panel(
    *,
    universe: str = "sp500",
    yahoo_period: str,
    start: str,
    end: str,
    max_tickers: int,
    refresh_cache: bool,
    pit_universe: bool,
) -> dict[str, pd.DataFrame]:
    import universe_scanner as us  # type: ignore[import-not-found]

    years = _parse_daily_period_years(yahoo_period)
    end_ts = pd.Timestamp(end) if str(end).strip() else pd.Timestamp.today()
    start_ts = pd.Timestamp(start)
    u = universe.lower().strip()

    if u in ("russell3000", "r3k", "russell_3000"):
        tickers = fetch_russell3000_tickers()
        if max_tickers and max_tickers > 0:
            tickers = tickers[: int(max_tickers)]
        label = "Russell 3000 snapshot"
    elif pit_universe:
        tickers = sp500_union_tickers(start_ts, end_ts)
        if max_tickers and max_tickers > 0:
            snap = fetch_sp500_constituents_with_sectors()["ticker"].astype(str).tolist()
            tickers = [t for t in snap if t in set(tickers)][: int(max_tickers)]
        label = "PIT S&P 500 union"
    else:
        tab = fetch_sp500_constituents_with_sectors()
        tickers = tab["ticker"].astype(str).tolist()
        if max_tickers and max_tickers > 0:
            tickers = tickers[: int(max_tickers)]
        label = "S&P 500 snapshot"

    print(f"🔍 Downloading {label} ({len(tickers)} names) …", flush=True)
    raw = us.download_and_cache_data(
        tickers,
        timeframe="1d",
        lookback_years=years,
        max_age_hours=0.0 if refresh_cache else 24.0,
    )
    equity_dict: dict[str, pd.DataFrame] = {}
    for t, df in raw.items():
        if df is None or df.empty:
            continue
        df2 = df.copy()
        df2.index = pd.to_datetime(df2.index).tz_localize(None)
        df2 = df2.sort_index()
        out = _compute_daily_backtest_features(df2)
        for vol_col in ("volume", "Volume"):
            if vol_col in df2.columns:
                out["volume"] = df2[vol_col].astype(np.float64)
                break
        equity_dict[t] = out
    return equity_dict


def build_engine(
    *,
    equity_dict: dict[str, pd.DataFrame],
    cfg: Sp500MomentumConfig,
    start: str,
    end: str,
    yahoo_period: str,
    refresh_shares: bool,
    pit_universe: bool = True,
) -> Sp500MomentumIndex:
    tickers = sorted(equity_dict.keys())
    sector_tab = fetch_sp500_constituents_with_sectors()
    sector_map = dict(zip(sector_tab["ticker"].astype(str), sector_tab["sector"].astype(str)))
    sector_map.update(_wiki_sector_map_for_tickers(set(tickers)))

    spy_df = _compute_daily_backtest_features(
        DataLoader().fetch_daily("SPY", period=yahoo_period)
    )
    eng = Sp500MomentumIndex(config=cfg, sector_map=sector_map, spy_panel=spy_df)

    close_pan = []
    for t in tickers:
        s = equity_dict[t]["close"].astype(np.float64)
        s.index = pd.to_datetime(equity_dict[t].index).tz_localize(None)
        close_pan.append(s.rename(t))
    close_df = pd.concat(close_pan, axis=1, sort=True).sort_index()

    if cfg.cap_proxy == "mkt_cap":
        shares = load_shares_outstanding_dict(
            tickers, start="1995-01-01", refresh=bool(refresh_shares)
        )
        eng.market_cap_df = build_market_cap_panel(close_df, shares)

    end_ts = pd.Timestamp(end) if str(end).strip() else close_df.index.max()
    if pit_universe:
        tmp = close_df.pct_change().rolling(231).std()
        bm = tmp.resample("BME").last().index
        if cfg.rebalance == "semiannual":
            bm = bm[bm.month.isin(tuple(cfg.semiannual_months))]
        bm = bm[(bm >= pd.Timestamp(start)) & (bm <= end_ts)]
        try:
            eng.membership_mask = sp500_membership_mask(tickers, bm)
        except ImportError:
            pass
    return eng


def yearly_stats(r: pd.Series) -> pd.DataFrame:
    rows: list[dict] = []
    for year, g in r.groupby(r.index.year):
        eq = (1.0 + g).cumprod()
        rows.append({
            "year": int(year),
            "return_pct": float(eq.iloc[-1] - 1.0) * 100.0,
            "max_dd_pct": float((eq / eq.cummax() - 1.0).min()) * 100.0,
            "n_days": len(g),
        })
    return pd.DataFrame(rows)


def summarize_returns(
    r: pd.Series,
    *,
    capital: float,
    spy_df: pd.DataFrame,
) -> dict[str, Any]:
    r = r.sort_index().astype(np.float64)
    cap = float(capital)
    eq = cap * (1.0 + r).cumprod()
    n = len(r)
    years = n / 252.0
    end_eq = float(eq.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 / eq.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_r = spy_df["ret"].astype(float)
    spy_r.index = pd.to_datetime(spy_r.index).tz_localize(None)
    aligned = pd.DataFrame({"mom": r, "spy": spy_r.reindex(r.index).fillna(0.0)}).dropna()
    rho = float(aligned.corr().iloc[0, 1]) if len(aligned) > 2 else float("nan")
    beta = float("nan")
    if len(aligned) > 10:
        sv = float(aligned["spy"].var())
        if sv > 1e-14:
            beta = float(aligned["mom"].cov(aligned["spy"]) / sv)
    return {
        "n_trading_days": int(n),
        "total_return_pct": round(100.0 * tot, 4),
        "cagr_pct": round(100.0 * cagr, 4),
        "sharpe_daily": round(sharpe, 4),
        "max_drawdown_pct": round(100.0 * max_dd, 4),
        "corr_vs_spy_daily": round(rho, 4),
        "beta_vs_spy": round(beta, 4),
        "start": str(r.index.min().date()),
        "end": str(r.index.max().date()),
    }


def run_backtest(
    *,
    equity_dict: dict[str, pd.DataFrame],
    cfg: Sp500MomentumConfig,
    start: str,
    end: str,
    yahoo_period: str,
    capital: float,
    refresh_shares: bool,
    pit_universe: bool = True,
    verbose: bool = True,
) -> tuple[pd.Series, pd.DataFrame, dict[str, Any]]:
    eng = build_engine(
        equity_dict=equity_dict,
        cfg=cfg,
        start=start,
        end=end,
        yahoo_period=yahoo_period,
        refresh_shares=refresh_shares,
        pit_universe=pit_universe,
    )
    if verbose:
        print(f"Running backtest on {len(equity_dict)} names …", flush=True)
    daily_ret = eng.generate_returns(equity_dict, verbose=verbose)
    rebal = eng.generate_rebalance_log(equity_dict)

    daily_ret = daily_ret.sort_index()
    daily_ret.index = pd.to_datetime(daily_ret.index).tz_localize(None)
    mask = daily_ret.index >= pd.Timestamp(start)
    if str(end).strip():
        mask &= daily_ret.index <= pd.Timestamp(end)
    r = daily_ret.loc[mask].fillna(0.0).astype(np.float64)

    spy_df = _compute_daily_backtest_features(
        DataLoader().fetch_daily("SPY", period=yahoo_period)
    )
    meta = summarize_returns(r, capital=capital, spy_df=spy_df)
    meta["config"] = {k: getattr(cfg, k) for k in cfg.__dataclass_fields__}
    return r, rebal, meta


def write_run_artifacts(
    *,
    prefix: Path,
    slug: str,
    r: pd.Series,
    rebal: pd.DataFrame,
    meta: dict[str, Any],
    capital: float,
    command: str,
) -> dict[str, str]:
    prefix = prefix.expanduser().resolve()
    prefix.parent.mkdir(parents=True, exist_ok=True)
    daily_path = Path(f"{prefix}_{slug}_daily.csv")
    rebal_path = Path(f"{prefix}_{slug}_rebalances.csv")
    yearly_path = Path(f"{prefix}_{slug}_yearly.csv")
    meta_path = Path(f"{prefix}_{slug}_meta.json")
    metrics_path = Path(f"{prefix}_{slug}_metrics.txt")

    cap = float(capital)
    pnl = r * cap
    eq_unit = (1.0 + r).cumprod()
    eq_usd = cap * eq_unit
    pd.DataFrame({
        "date": r.index.strftime("%Y-%m-%d"),
        "daily_ret": r.values,
        "daily_pnl_usd": pnl.values,
        "equity_unit": eq_unit.values,
        "equity_usd": eq_usd.values,
    }).to_csv(daily_path, index=False)
    yearly_stats(r).to_csv(yearly_path, index=False)
    if len(rebal):
        rebal.to_csv(rebal_path, index=False)
    else:
        pd.DataFrame().to_csv(rebal_path, index=False)

    meta_out = dict(meta)
    meta_out["command"] = command
    meta_out["daily_csv"] = str(daily_path)
    meta_out["rebalances_csv"] = str(rebal_path)
    meta_out["yearly_csv"] = str(yearly_path)
    meta_path.write_text(json.dumps(meta_out, indent=2, default=str) + "\n", encoding="utf-8")
    metrics_path.write_text(
        "\n".join([
            command, "",
            f"Window: {meta['start']} → {meta['end']}",
            f"Return: {meta['total_return_pct']:+.2f}%  CAGR: {meta['cagr_pct']:+.2f}%",
            f"Sharpe: {meta['sharpe_daily']:.3f}  Max DD: {meta['max_drawdown_pct']:.2f}%",
            f"β(SPY): {meta['beta_vs_spy']:.3f}  ρ(SPY): {meta['corr_vs_spy_daily']:.3f}",
        ]) + "\n",
        encoding="utf-8",
    )
    return {
        "daily": str(daily_path),
        "rebalances": str(rebal_path),
        "yearly": str(yearly_path),
        "meta": str(meta_path),
        "metrics": str(metrics_path),
    }
