#!/usr/bin/env python3
"""
Cross-sectional dollar-neutral long–short backtest (RenTech / stat-arb scaffold).

  Long top --top-n by signal, short bottom --bottom-n, equal weight per leg.
  PnL uses realized forward return from daily_sequence y_* (same as 6_* pipeline).

Signal sources:
  parquet — merge --predictions (e.g. xgb_val.parquet) on symbol+asof_date, column --signal-col
  random  — i.i.d. normal (sanity check for plumbing only)

Example:
  .venv/bin/python RenTech/run_cs_neutral_backtest.py \\
    --prefix daily_sequence --hold 1 \\
    --signal-source parquet --predictions xgb_val.parquet --signal-col pred_cc_reg \\
    --top-n 25 --bottom-n 25 --out-daily rentech_cs_daily.csv
"""

from __future__ import annotations

import argparse
import os
import sys

# Allow `python RenTech/run_cs_neutral_backtest.py` from repo root
_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), ".."))
if _ROOT not in sys.path:
    sys.path.insert(0, _ROOT)

from RenTech.stat_arb import (
    apply_round_trip_cost,
    attach_signal,
    cross_sectional_long_short_daily,
    load_forward_return_panel,
    summarize_daily_returns,
)


def main() -> None:
    p = argparse.ArgumentParser(description=__doc__)
    p.add_argument("--prefix", default="daily_sequence", help="daily_sequence_* prefix")
    p.add_argument("--hold", type=int, choices=(1, 3, 5), default=1, help="Forward return horizon")
    p.add_argument(
        "--signal-source",
        choices=("parquet", "random"),
        default="parquet",
        help="Where signal comes from",
    )
    p.add_argument("--predictions", default="xgb_val.parquet", help="For parquet source")
    p.add_argument("--signal-col", default="pred_cc_reg", dest="signal_col")
    p.add_argument("--random-seed", type=int, default=42, dest="random_seed")
    p.add_argument("--top-n", type=int, default=25, dest="top_n")
    p.add_argument("--bottom-n", type=int, default=25, dest="bottom_n")
    p.add_argument("--min-names", type=int, default=80, help="Min eligible names per day")
    p.add_argument("--round-trip-bps", type=float, default=8.0, help="Cost per day (long+short legs)")
    p.add_argument("--out-daily", default="rentech_cs_neutral_daily.csv")
    args = p.parse_args()

    os.chdir(_ROOT)

    panel = load_forward_return_panel(args.prefix, hold_days=args.hold)  # type: ignore[arg-type]
    print(f"📌 Panel rows: {len(panel):,} | dates: {panel['asof_date'].nunique():,}")

    try:
        with_signal = attach_signal(
            panel,
            source=args.signal_source,
            predictions_path=args.predictions if args.signal_source == "parquet" else None,
            signal_col=args.signal_col,
            random_seed=args.random_seed,
        )
    except FileNotFoundError as e:
        print(f"❌ {e}", file=sys.stderr)
        if args.signal_source == "parquet":
            print(
                "   Build scores first, e.g.:\n"
                "   .venv/bin/python 6_predict_xgb_daily.py \\\n"
                "     --manifest xgb_daily_models/manifest.json --split val --out xgb_val.parquet",
                file=sys.stderr,
            )
        sys.exit(1)

    print(f"📌 Rows after signal merge: {len(with_signal):,}")

    daily = cross_sectional_long_short_daily(
        with_signal,
        top_n=args.top_n,
        bottom_n=args.bottom_n,
        min_names=args.min_names,
    )
    if daily.empty:
        print("❌ No days passed min-names / portfolio filters.", file=sys.stderr)
        sys.exit(1)

    daily["port_ret_net"] = apply_round_trip_cost(daily["port_ret_gross"], args.round_trip_bps)
    daily.to_csv(args.out_daily, index=False)

    st = summarize_daily_returns(daily["port_ret_net"])
    print("\n=== Dollar-neutral L/S (net of round-trip bps) ===")
    for k, v in st.items():
        print(f"  {k}: {v}")
    print(f"\n💾 Daily series → {args.out_daily}")


if __name__ == "__main__":
    main()
