#!/usr/bin/env python3
"""
Equal-capital sleeves: combine chosen **D** strategies into one portfolio equity curve.

Each sleeve uses ``capital / N`` nominal starting equity, trades **independently** (same
engine as ``runner.py``: no overlapping trades *within* a sleeve). Aggregate equity is the
sum of sleeve equity curves (aligned on Theta × SPY overlap days).

Default basket is the **top-9 Sharpe** list from ``diverse_theta_strategies_v1_2021_2024``
(D039, D018, D095, D081, D057, D065, D041, D046, D022).

Run::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 \\
        .venv/bin/python -m RenTech.strategy_stack.diverse_theta_strategies_v1.portfolio_top_strategies \\
        --start 2016-01-04 --end 2026-12-31 \\
        --capital 1000000 \\
        --out-csv RenTech/data/logs/diverse_theta_top9_portfolio_equity.csv \\
        --out-json RenTech/data/logs/diverse_theta_top9_portfolio_summary.json

``--end`` may clip to latest Theta chunk; use empty ``--end`` to use all overlapping data through chunk max.
"""
from __future__ import annotations

import argparse
import json
import math
import sys
import time
from pathlib import Path
from typing import Any, Callable

import numpy as np
import pandas as pd

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

from RenTech.core.options_data_loader import OptionChain
from RenTech.strategy_stack import research_literature_theta_strategies as L
from RenTech.strategy_stack.diverse_theta_strategies_v1.context import ResearchContext
from RenTech.strategy_stack.diverse_theta_strategies_v1.panel import augment_research_panel
from RenTech.strategy_stack.diverse_theta_strategies_v1.runner import (  # reuse trade dispatch + loader
    _load_strategy_module,
    _trade_fn,
)

DEFAULT_TOP_9_SID = ("D039", "D018", "D095", "D081", "D057", "D065", "D041", "D046", "D022")

SignalFn = Callable[[int, pd.Series, OptionChain, float], bool]


def sid_to_module_index(sid: str) -> int:
    s = sid.strip().upper()
    if not s.startswith("D") or len(s) != 4:
        raise ValueError(f"Expected sid like D018, got {sid!r}")
    return int(s[1:])


def _cagr_decimal(start_pv: float, end_pv: float, years: float) -> float | None:
    if start_pv <= 0 or years <= 0:
        return None
    if end_pv <= 0:
        return None
    return (end_pv / start_pv) ** (1.0 / years) - 1.0


def main() -> None:
    ap = argparse.ArgumentParser(description="Equal-weight sleeve portfolio from D-catalog modules")
    ap.add_argument("--theta-dir", type=Path, default=L._DEFAULT_THETA)
    ap.add_argument("--capital", type=float, default=1_000_000.0)
    ap.add_argument("--start", type=str, default="2016-01-04")
    ap.add_argument("--end", type=str, default="", help="YYYY-MM-DD inclusive; omit for max theta")
    ap.add_argument("--sids", type=str, nargs="*", default=list(DEFAULT_TOP_9_SID))
    ap.add_argument(
        "--out-csv",
        type=Path,
        default=_REPO_ROOT / "RenTech" / "data" / "logs" / "diverse_theta_top9_portfolio_equity.csv",
    )
    ap.add_argument(
        "--out-json",
        type=Path,
        default=_REPO_ROOT / "RenTech" / "data" / "logs" / "diverse_theta_top9_portfolio_summary.json",
    )
    args = ap.parse_args()

    sids: list[str] = [str(x).strip().upper() for x in args.sids]
    n = len(sids)
    if n == 0:
        raise SystemExit("Need at least one --sids")

    sleeve_cap = float(args.capital) / float(n)
    theta_dir = Path(args.theta_dir).expanduser().resolve()

    t0 = time.perf_counter()
    days, panel0, get_chain, iv_atm, skew, n_contracts, spy_wide = L.prepare_theta_research_context(
        theta_dir=theta_dir,
        capital=float(args.capital),
        start=str(args.start).strip(),
        end=str(args.end).strip(),
        max_days=0,
    )
    panel = augment_research_panel(panel0)
    ctx = ResearchContext(days, panel, get_chain, iv_atm, skew, n_contracts)
    idx = pd.DatetimeIndex([L._norm(d) for d in days])
    n_sess = len(days)
    calendar_years = (idx[-1] - idx[0]).days / 365.25 if n_sess >= 2 else float("nan")
    trading_years = n_sess / 252.0 if n_sess >= 2 else float("nan")
    print(
        f"Context: {n_sess} sessions from {idx[0].date()} to {idx[-1].date()} "
        f"({time.perf_counter() - t0:.0f}s precompute elapsed)",
        flush=True,
    )

    sleeve_equity = pd.DataFrame(index=idx, dtype=float)
    sleeve_summary: list[dict[str, Any]] = []
    cumulative_pnl_components: dict[str, float] = {}

    for sid in sids:
        k = sid_to_module_index(sid)
        mod = _load_strategy_module(k)
        meta = dict(getattr(mod, "META"))
        sid_meta = str(meta.get("sid", sid))
        hold = int(getattr(mod, "HOLD_SESSIONS"))
        tk = str(getattr(mod, "TRADE_KIND"))
        tp = tuple(getattr(mod, "TRADE_PARAMS"))
        tfn = _trade_fn(tk, tp)
        wants = getattr(mod, "wants_entry")

        def _sig(
            i: int,
            row: pd.Series,
            ch: OptionChain,
            spy: float,
            _w=wants,
        ) -> bool:
            return bool(_w(i, row, ch, spy, ctx))

        ex, pnls, ntr = L.run_signal_backtest(days, get_chain, panel, _sig, hold, tfn, tp)
        eq, sh = L.equity_curve_from_realized(ex, pnls, days, sleeve_cap)
        sleeve_equity[sid_meta] = eq.reindex(idx).astype(float)
        sleeve_pnl = float(eq.iloc[-1]) - sleeve_cap if len(eq) else 0.0
        cumulative_pnl_components[sid_meta] = sleeve_pnl
        sleeve_summary.append(
            {
                "sid": sid_meta,
                "module": f"s{k:03d}",
                "title": str(meta.get("title", "")),
                "trade_kind": tk,
                "trade_params": list(tp),
                "hold_sessions": hold,
                "trades": int(ntr),
                "sleeve_capital_start_usd": sleeve_cap,
                "sleeve_equity_end_usd": float(eq.iloc[-1]) if len(eq) else sleeve_cap,
                "sleeve_total_pnl_usd": sleeve_pnl,
                "sharpe_daily_returns": float(sh) if math.isfinite(sh) else None,
            }
        )

    port_eq = sleeve_equity.sum(axis=1)
    port_ret = port_eq.pct_change().replace([np.inf, -np.inf], np.nan).fillna(0.0)
    sharpe_ptf = float(L.sharpe_daily_returns(port_ret)) if n_sess >= 50 else float("nan")

    start_cap = float(args.capital)
    end_cap = float(port_eq.iloc[-1]) if len(port_eq) else start_cap
    total_return_pct = (end_cap / start_cap - 1.0) * 100.0
    cum_pnl = end_cap - start_cap

    cagr_cal = (
        _cagr_decimal(start_cap, end_cap, float(calendar_years))
        if math.isfinite(calendar_years) and calendar_years > 0
        else None
    )
    cagr_trade = (
        _cagr_decimal(start_cap, end_cap, float(trading_years))
        if math.isfinite(trading_years) and trading_years > 0
        else None
    )

    summary: dict[str, Any] = {
        "capital_start_usd": start_cap,
        "capital_end_usd": end_cap,
        "total_return_pct": total_return_pct,
        "total_pnl_usd": cum_pnl,
        "sessions": int(n_sess),
        "calendar_years_span": calendar_years,
        "approx_trading_years_sess_over_252": trading_years,
        "cagr_decimal_calendar_span": None if cagr_cal is None else float(cagr_cal),
        "cagr_pct_calendar_span": None if cagr_cal is None else float(cagr_cal) * 100.0,
        "cagr_decimal_tradingyears_252": None if cagr_trade is None else float(cagr_trade),
        "cagr_pct_tradingyears_252": None if cagr_trade is None else float(cagr_trade) * 100.0,
        "portfolio_sharpe_daily_returns": sharpe_ptf if math.isfinite(sharpe_ptf) else None,
        "sleeves_equal_weight_n": int(n),
        "sleeve_capital_each_usd": sleeve_cap,
        "sleeves_ordered": list(sids),
        "theta_dir": str(theta_dir),
        "start_session": str(idx[0].date()) if len(idx) else "",
        "last_session": str(idx[-1].date()) if len(idx) else "",
        "component_total_pnl_by_sid_usd": cumulative_pnl_components,
        "component_sleeves": sleeve_summary,
    }

    out_csv = Path(args.out_csv).expanduser().resolve()
    out_csv.parent.mkdir(parents=True, exist_ok=True)
    out_sheet = sleeve_equity.copy()
    out_sheet.insert(0, "portfolio_equity_usd", port_eq)
    out_sheet.to_csv(out_csv)

    out_json = Path(args.out_json).expanduser().resolve()
    out_json.parent.mkdir(parents=True, exist_ok=True)
    out_json.write_text(json.dumps(summary, indent=2), encoding="utf-8")

    print(json.dumps(summary, indent=2)[:4000])
    print(f"\nWrote CSV {out_csv}", flush=True)
    print(f"Wrote JSON {out_json}", flush=True)


if __name__ == "__main__":
    main()
