#!/usr/bin/env python3
"""Trade-level stats for intraday MA-slope (confirm_entry_4b vs baseline)."""

from __future__ import annotations

import argparse
import json
import sys
from dataclasses import replace
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))

from RenTech.strategy_stack.alpaca_minute_loader import (
    DEFAULT_ALPACA_RTH_DIR,
    compound_intraday_to_daily,
    list_parquet_symbols,
    load_equity_panels,
)
from RenTech.strategy_stack.ma_slope_intraday_enhanced import (
    EnhancedIntradayEngine,
    baseline_enhanced_config,
    metrics_daily,
)
from RenTech.strategy_stack.multi_strategy_manager import _align_panel_frames
from RenTech.strategy_stack.run_johansen_triplet_sp500 import load_sp500_sectors

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


def _session_key(index: pd.DatetimeIndex) -> pd.Series:
    return pd.Series(index.normalize(), index=index)


def run_with_weights(engine: EnhancedIntradayEngine, intra: dict, top_n: int, panels) -> tuple[pd.Series, pd.DataFrame, pd.DataFrame]:
    """Return (port_r, exec_w, ret_df)."""
    ret_pan = []
    for t, df in sorted(intra.items()):
        idx = pd.to_datetime(df.index).tz_localize(None)
        dfx = df.copy()
        dfx.index = idx
        ret_pan.append(dfx["ret"].astype(np.float64).rename(t))
    ret_df, _ = _align_panel_frames(ret_pan)
    master = ret_df.index.sort_values()
    ret_df = ret_df.reindex(master).fillna(0.0)
    from RenTech.strategy_stack.ma_slope_intraday_enhanced import EnhancedPanels

    panels = EnhancedPanels(
        panels.score.reindex(master).reindex(columns=ret_df.columns),
        panels.fast_slope.reindex(master).reindex(columns=ret_df.columns),
        panels.fast_ma.reindex(master).reindex(columns=ret_df.columns),
        panels.close.reindex(master).reindex(columns=ret_df.columns),
        panels.vwap.reindex(master).reindex(columns=ret_df.columns),
        panels.or_high.reindex(master).reindex(columns=ret_df.columns),
        panels.sector,
    )
    target_w = engine.target_weights_enhanced(panels, top_n)
    target_w = engine._apply_momentum_exits(target_w, panels)
    exec_w = target_w.shift(1).fillna(0.0)
    exec_w.loc[engine._first_bar_mask(master)] = 0.0
    port_r = pd.Series((exec_w.to_numpy() * ret_df.to_numpy()).sum(axis=1), index=master)
    port_r = engine._apply_session_stop(port_r, exec_w)
    port_r = engine._apply_slippage(port_r, exec_w)
    return port_r, exec_w, ret_df


def session_stats(port_r: pd.Series, exec_w: pd.DataFrame) -> dict:
    daily = compound_intraday_to_daily(port_r).dropna()
    ses = _session_key(pd.DatetimeIndex(exec_w.index))
    gross = exec_w.sum(axis=1)
    daily_gross = gross.groupby(ses).max()

    traded = daily[daily_gross.reindex(daily.index).fillna(0) > 1e-9]
    flat = daily[daily_gross.reindex(daily.index).fillna(0) <= 1e-9]
    wins = traded[traded > 0]
    losses = traded[traded < 0]

    def _avg(x: pd.Series) -> float:
        return float(x.mean() * 100) if len(x) else float("nan")

    pf_num = float(wins.sum()) if len(wins) else 0.0
    pf_den = float(-losses.sum()) if len(losses) else 0.0
    return {
        "n_calendar_days": int(len(daily)),
        "n_sessions_traded": int(len(traded)),
        "n_sessions_flat": int(len(flat)),
        "pct_days_traded": float(len(traded) / len(daily) * 100) if len(daily) else 0.0,
        "session_win_rate_pct": float(len(wins) / len(traded) * 100) if len(traded) else 0.0,
        "session_avg_return_pct": _avg(traded),
        "session_avg_winner_pct": _avg(wins),
        "session_avg_loser_pct": _avg(losses),
        "session_median_return_pct": float(traded.median() * 100) if len(traded) else float("nan"),
        "session_best_pct": float(traded.max() * 100) if len(traded) else float("nan"),
        "session_worst_pct": float(traded.min() * 100) if len(traded) else float("nan"),
        "profit_factor": float(pf_num / pf_den) if pf_den > 1e-12 else float("inf"),
        "avg_gross_exposure": float(daily_gross[daily_gross > 0].mean()) if (daily_gross > 0).any() else 0.0,
    }


def name_trade_stats(exec_w: pd.DataFrame, ret_df: pd.DataFrame) -> dict:
    """One trade = one symbol held for a session (weight > 0)."""
    ses = _session_key(pd.DatetimeIndex(exec_w.index))
    rows = []
    for sym in exec_w.columns:
        w = exec_w[sym].astype(float)
        r = ret_df[sym].astype(float)
        for day, idx in w.groupby(ses).groups.items():
            wi = w.loc[idx]
            if float(wi.max()) <= 1e-9:
                continue
            # Constant weight within session after entry; compound bar returns.
            ri = r.loc[idx]
            active = wi > 1e-9
            if not active.any():
                continue
            tr = float(np.prod(1.0 + ri.values[active.values]) - 1.0)
            rows.append({"date": pd.Timestamp(day), "symbol": sym, "return_pct": tr * 100, "weight": float(wi.max())})
    if not rows:
        return {}
    df = pd.DataFrame(rows)
    wins = df[df["return_pct"] > 0]
    losses = df[df["return_pct"] < 0]
    flat = df[df["return_pct"] == 0]
    pf_num = float(wins["return_pct"].sum()) if len(wins) else 0.0
    pf_den = float(-losses["return_pct"].sum()) if len(losses) else 0.0
    return {
        "n_name_trades": int(len(df)),
        "n_unique_symbols": int(df["symbol"].nunique()),
        "avg_names_per_session": float(df.groupby("date").size().mean()),
        "name_win_rate_pct": float(len(wins) / len(df) * 100),
        "name_avg_return_pct": float(df["return_pct"].mean()),
        "name_avg_winner_pct": float(wins["return_pct"].mean()) if len(wins) else float("nan"),
        "name_avg_loser_pct": float(losses["return_pct"].mean()) if len(losses) else float("nan"),
        "name_median_return_pct": float(df["return_pct"].median()),
        "name_best_pct": float(df["return_pct"].max()),
        "name_worst_pct": float(df["return_pct"].min()),
        "name_profit_factor": float(pf_num / pf_den) if pf_den > 1e-12 else float("inf"),
        "n_flat_name_trades": int(len(flat)),
    }


def yearly_session_win_rate(port_r: pd.Series, exec_w: pd.DataFrame) -> list[dict]:
    daily = compound_intraday_to_daily(port_r).dropna()
    ses = _session_key(pd.DatetimeIndex(exec_w.index))
    gross = exec_w.sum(axis=1)
    daily_gross = gross.groupby(ses).max()
    traded = daily[daily_gross.reindex(daily.index).fillna(0) > 1e-9]
    out = []
    for y, grp in traded.groupby(traded.index.year):
        wins = (grp > 0).sum()
        out.append(
            {
                "year": int(y),
                "n_sessions": int(len(grp)),
                "win_rate_pct": float(wins / len(grp) * 100) if len(grp) else 0.0,
                "avg_return_pct": float(grp.mean() * 100),
            }
        )
    return out


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__)
    ap.add_argument("--start", default="2020-01-02")
    ap.add_argument("--end", default="2026-06-26")
    ap.add_argument("--top-n", type=int, default=10)
    ap.add_argument("--max-tickers", type=int, default=500)
    ap.add_argument("--out-prefix", type=Path, default=LOGS / "ma_slope_intraday_trade_stats")
    args = ap.parse_args()

    syms = list_parquet_symbols(DEFAULT_ALPACA_RTH_DIR)[: int(args.max_tickers)]
    if "SPY" not in syms:
        syms.insert(0, "SPY")
    sector_map = dict(
        zip(
            load_sp500_sectors()["ticker"].astype(str).str.upper(),
            load_sp500_sectors()["sector"].astype(str),
        )
    )

    print("Loading panels …", flush=True)
    intra, _ = load_equity_panels(
        syms,
        start="2019-11-01",
        end=str(args.end),
        warmup_sessions=15,
    )
    ret_start = pd.Timestamp(args.start)
    end = pd.Timestamp(args.end)

    variants = [
        ("baseline", baseline_enhanced_config()),
        ("confirm_entry_4b", replace(baseline_enhanced_config(), hold_mode="confirm_entry", confirm_lag_bars=4)),
        (
            "confirm_4b_5bps",
            replace(baseline_enhanced_config(), hold_mode="confirm_entry", confirm_lag_bars=4, slippage_bps=5.0),
        ),
    ]

    out: dict = {"start": args.start, "end": args.end, "n_loaded": len(intra), "variants": {}}
    for name, cfg in variants:
        print(f"  {name} …", flush=True)
        eng = EnhancedIntradayEngine(config=cfg)
        panels = eng.build_panels(intra, sector_map)
        port_r, exec_w, ret_df = run_with_weights(eng, intra, int(args.top_n), panels)
        port_r = port_r.loc[(port_r.index >= ret_start) & (port_r.index <= end)]
        exec_w = exec_w.loc[(exec_w.index >= ret_start) & (exec_w.index <= end)]
        ret_df = ret_df.loc[(ret_df.index >= ret_start) & (ret_df.index <= end)]
        m = metrics_daily(port_r)
        ss = session_stats(port_r, exec_w)
        nt = name_trade_stats(exec_w, ret_df)
        yr = yearly_session_win_rate(port_r, exec_w)
        out["variants"][name] = {
            "headline": m,
            "session": ss,
            "name_trades": nt,
            "yearly_session": yr,
        }

    prefix = args.out_prefix.expanduser().resolve()
    json_path = Path(f"{prefix}_top{args.top_n}.json")
    json_path.write_text(json.dumps(out, indent=2) + "\n")

    # print summary table
    for name, v in out["variants"].items():
        h, s, n = v["headline"], v["session"], v["name_trades"]
        print(f"\n=== {name} ===", flush=True)
        print(
            f"  Return {h.get('total_return_pct', 0):+.1f}%  Sharpe {h.get('sharpe', 0):.2f}  "
            f"MaxDD {h.get('max_dd_pct', 0):.1f}%",
            flush=True,
        )
        print(
            f"  Sessions traded {s['n_sessions_traded']} ({s['pct_days_traded']:.0f}% of days)  "
            f"win rate {s['session_win_rate_pct']:.1f}%  profit factor {s['profit_factor']:.2f}",
            flush=True,
        )
        print(
            f"  Avg session {s['session_avg_return_pct']:+.3f}%  "
            f"avg winner {s['session_avg_winner_pct']:+.3f}%  avg loser {s['session_avg_loser_pct']:+.3f}%",
            flush=True,
        )
        print(
            f"  Name-trades {n['n_name_trades']}  name win rate {n['name_win_rate_pct']:.1f}%  "
            f"avg name {n['name_avg_return_pct']:+.3f}%  (+{n['name_avg_winner_pct']:.3f}% / {n['name_avg_loser_pct']:.3f}%)",
            flush=True,
        )
    print(f"\nWrote {json_path}", flush=True)


if __name__ == "__main__":
    main()
