#!/usr/bin/env python3
"""
Build a reproducible multi-sleeve literature portfolio and export:

- daily portfolio PnL / equity
- detailed trade log (entry/exit, legs JSON, qty, per-trade margin, PnL)
- summary JSON

Sizing modes:
- ``capital_scaled`` (default): per-sleeve capital buckets with integer contracts
- ``fixed_weight``: legacy one-lot sleeve blend weighted by ``1/N``
"""
from __future__ import annotations

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

import numpy as np
import pandas as pd

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

from RenTech.strategy_stack import research_literature_theta_strategies as L
from RenTech.strategy_stack.literature_search_agent import _compile_signal, _compile_trade
from RenTech.strategy_stack.literature_strategy_catalog import StrategySpec, build_catalog_100

_DEFAULT_THETA = _REPO_ROOT / "RenTech" / "data" / "theta_chunks"
_DEFAULT_JSON = (
    _REPO_ROOT
    / "RenTech"
    / "data"
    / "logs"
    / "literature_search_agent_batch_2016_2022_v3_diverse_caps.json"
)
_OUT_DAILY = _REPO_ROOT / "RenTech" / "data" / "logs" / "literature_low_corr_portfolio_daily_pnl.csv"
_OUT_TRADES = _REPO_ROOT / "RenTech" / "data" / "logs" / "literature_low_corr_portfolio_trade_log.csv"


def _catalog_by_sid() -> dict[str, StrategySpec]:
    return {s.sid: s for s in build_catalog_100()}


def _repo_rel(p: Path) -> str:
    r = p.resolve()
    try:
        return str(r.relative_to(_REPO_ROOT))
    except ValueError:
        return str(r)


def _resolve_sids(from_json: Path | None, selection_key: str, sids_csv: str) -> list[str]:
    if str(sids_csv).strip():
        out = [x.strip() for x in str(sids_csv).split(",") if x.strip()]
        if not out:
            raise SystemExit("--sids was empty after parsing")
        return out
    pj = Path(from_json).expanduser().resolve() if from_json is not None else _DEFAULT_JSON
    if not pj.is_file():
        raise SystemExit(f"Batch JSON not found: {pj} (pass --from-json or --sids)")
    payload = json.loads(pj.read_text(encoding="utf-8"))
    sel = payload.get("selection") or {}
    key = str(selection_key).strip().lower()
    if key == "capped":
        sids = list(sel.get("greedy_capped") or [])
    elif key == "uncapped":
        sids = list(sel.get("greedy_uncapped") or [])
    else:
        raise SystemExit("--selection must be capped or uncapped")
    if not sids:
        alt = sel.get("greedy_uncapped" if key == "capped" else "greedy_capped") or []
        sids = list(alt)
    if not sids:
        raise SystemExit(f"No strategy ids in JSON selection ({selection_key})")
    return sids


def _margin_per_contract(
    trade_kind: str,
    entry_spy: float,
    legs: list[dict[str, Any]],
    regt_short_put_mult: float,
) -> float:
    shorts = [l for l in legs if str(l.get("position")) == "short"]
    longs = [l for l in legs if str(l.get("position")) == "long"]

    if trade_kind in ("vert", "vtc") and shorts and longs:
        w = abs(float(shorts[0]["strike"]) - float(longs[0]["strike"]))
        return max(1.0, 100.0 * w)

    if trade_kind == "sl":
        debit = sum(float(l.get("entry_ask") or 0.0) for l in longs) * 100.0
        return max(1.0, debit)

    if trade_kind in ("put", "rr") and shorts:
        return max(1.0, 100.0 * float(shorts[0]["strike"]) * float(regt_short_put_mult))

    if trade_kind in ("ss", "sg") and shorts:
        credit = sum(float(l.get("entry_bid") or 0.0) for l in shorts) * 100.0
        base = max(0.20 * float(entry_spy) * 100.0, 0.10 * float(entry_spy) * 100.0 + credit)
        return max(1.0, base)

    return max(1.0, 0.20 * float(entry_spy) * 100.0)


def main() -> None:
    ap = argparse.ArgumentParser(description="Capital-scaled low-correlation literature portfolio export")
    ap.add_argument("--theta-dir", type=Path, default=_DEFAULT_THETA)
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--start", type=str, default="2016-01-04")
    ap.add_argument("--end", type=str, default="2022-12-31")
    ap.add_argument("--max-days", type=int, default=0)
    ap.add_argument(
        "--from-json",
        type=Path,
        default=None,
        help=f"Batch JSON from literature_search_agent (default: {_DEFAULT_JSON.name})",
    )
    ap.add_argument(
        "--selection",
        type=str,
        default="capped",
        choices=("capped", "uncapped"),
        help="Which greedy list to read from JSON when --sids is omitted",
    )
    ap.add_argument("--sids", type=str, default="", help="Comma-separated S000 ids; overrides JSON")
    ap.add_argument(
        "--sizing-mode",
        type=str,
        default="capital_scaled",
        choices=("capital_scaled", "fixed_weight"),
        help="capital_scaled is canonical; fixed_weight preserves old one-lot behavior",
    )
    ap.add_argument(
        "--regt-short-put-mult",
        type=float,
        default=1.0,
        help="Multiplier on strike*100 for short put / short put leg in RR margin model",
    )
    ap.add_argument("--out-daily", type=Path, default=_OUT_DAILY)
    ap.add_argument("--out-trades", type=Path, default=_OUT_TRADES)
    args = ap.parse_args()

    sids = _resolve_sids(args.from_json, args.selection, str(args.sids))
    n = len(sids)
    weight = 1.0 / float(n)
    sleeve_capital = float(args.capital) / float(n)
    catalog = _catalog_by_sid()
    for sid in sids:
        if sid not in catalog:
            raise SystemExit(f"Unknown sid {sid!r} (rebuild catalog or fix --sids)")

    t0 = time.perf_counter()
    print(f"Building Theta context… ({n} sleeves)", flush=True)
    days, panel, get_chain, iv_atm, skew_put_minus_call_iv, n_contracts, spy_wide = (
        L.prepare_theta_research_context(
            theta_dir=args.theta_dir,
            capital=float(args.capital),
            start=str(args.start),
            end=str(args.end),
            max_days=int(args.max_days),
        )
    )
    print(f"Context ready in {(time.perf_counter() - t0) / 60:.2f} min; {len(days)} sessions.", flush=True)

    idx = pd.DatetimeIndex([L._norm(d) for d in days])
    daily_by_sid: dict[str, pd.Series] = {}
    all_trades: list[dict[str, object]] = []
    sleeve_stats: dict[str, dict[str, float | int]] = {}

    for sid in sids:
        spec = catalog[sid]
        sig = _compile_signal(spec, panel, iv_atm, skew_put_minus_call_iv, n_contracts, spy_wide)
        tfn = _compile_trade(spec)
        t1 = time.perf_counter()
        trades = L.run_signal_backtest_trades(
            days,
            get_chain,
            panel,
            sig,
            int(spec.hold),
            tfn,
            spec.trade_params,
            spec.trade_kind,
            sid=sid,
            family=spec.family,
            description=spec.description,
        )
        print(f"  {sid} {spec.family}: {len(trades)} trades in {(time.perf_counter()-t1):.1f}s", flush=True)
        rows: list[dict[str, Any]] = []
        if str(args.sizing_mode) == "capital_scaled":
            eq_sid = float(sleeve_capital)
            qty_nonzero = 0
            qty_sum = 0
            for t in trades:
                legs = json.loads(str(t["legs_json"])) if str(t.get("legs_json", "")).strip() else []
                mpc = _margin_per_contract(
                    str(spec.trade_kind),
                    float(t["entry_spy"]),
                    legs,
                    float(args.regt_short_put_mult),
                )
                qty = int(max(0, math.floor(eq_sid / mpc)))
                pnl_1x = float(t["pnl_usd"])
                pnl_scaled = pnl_1x * qty
                eq_before = eq_sid
                eq_sid += pnl_scaled
                row = dict(t)
                row["qty"] = qty
                row["margin_per_contract_usd"] = mpc
                row["capital_allocated_sleeve_usd"] = sleeve_capital
                row["equity_sleeve_before_usd"] = eq_before
                row["equity_sleeve_after_usd"] = eq_sid
                row["pnl_1x_usd"] = pnl_1x
                row["pnl_sleeve_usd"] = pnl_scaled
                row["weight"] = weight
                row["pnl_portfolio_usd"] = pnl_scaled
                rows.append(row)
                if qty > 0:
                    qty_nonzero += 1
                    qty_sum += qty
            sleeve_stats[sid] = {
                "n_trades": int(len(rows)),
                "n_nonzero_qty": int(qty_nonzero),
                "avg_qty_nonzero": float(qty_sum / qty_nonzero) if qty_nonzero else 0.0,
                "equity_end_usd": float(eq_sid),
            }
        else:
            for t in trades:
                row = dict(t)
                row["qty"] = 1
                row["margin_per_contract_usd"] = float("nan")
                row["capital_allocated_sleeve_usd"] = sleeve_capital
                row["equity_sleeve_before_usd"] = float("nan")
                row["equity_sleeve_after_usd"] = float("nan")
                row["pnl_1x_usd"] = float(t["pnl_usd"])
                row["pnl_sleeve_usd"] = float(t["pnl_usd"]) * weight
                row["weight"] = weight
                row["pnl_portfolio_usd"] = float(t["pnl_usd"]) * weight
                rows.append(row)
            sleeve_stats[sid] = {
                "n_trades": int(len(rows)),
                "n_nonzero_qty": int(len(rows)),
                "avg_qty_nonzero": 1.0 if len(rows) else 0.0,
                "equity_end_usd": float("nan"),
            }

        all_trades.extend(rows)
        ex = [pd.Timestamp(r["exit_date"]) for r in rows]
        pn = [float(r["pnl_portfolio_usd"]) for r in rows]
        daily_by_sid[sid] = L.daily_pnl_series(ex, pn, days)

    pnl_port = pd.Series(0.0, index=idx)
    for sid in sids:
        pnl_port = pnl_port.add(daily_by_sid[sid], fill_value=0.0)
    cum = pnl_port.cumsum()
    eq = float(args.capital) + cum
    ret = eq.pct_change().replace([np.inf, -np.inf], np.nan).fillna(0.0)
    sh = L.sharpe_daily_returns(ret)
    dd = float((eq / eq.cummax() - 1.0).min()) if len(eq) else float("nan")

    daily_out = pd.DataFrame(
        {
            "pnl_portfolio_usd": pnl_port,
            "cumulative_pnl_usd": cum,
            "equity_usd": eq,
            "daily_return": ret,
        }
    )
    for sid in sids:
        daily_out[f"pnl_portfolio_{sid}"] = daily_by_sid[sid]

    out_d = args.out_daily.expanduser().resolve()
    out_t = args.out_trades.expanduser().resolve()
    out_d.parent.mkdir(parents=True, exist_ok=True)
    out_t.parent.mkdir(parents=True, exist_ok=True)
    daily_out.to_csv(out_d)
    trades_df = pd.DataFrame(all_trades)
    if not trades_df.empty:
        trades_df = trades_df.sort_values(["exit_date", "sid"]).reset_index(drop=True)
    trades_df.to_csv(out_t, index=False)

    meta = {
        "script": "literature_low_corr_portfolio.py",
        "sizing_mode": str(args.sizing_mode),
        "regt_short_put_mult": float(args.regt_short_put_mult),
        "n_sleeves": n,
        "weight_each": weight,
        "capital_per_sleeve_start_usd": sleeve_capital,
        "sids": sids,
        "selection_source": str(args.from_json or _DEFAULT_JSON) if not str(args.sids).strip() else "--sids",
        "first_day": str(days[0].date()) if days else "",
        "last_day": str(days[-1].date()) if days else "",
        "sessions": len(days),
        "capital_start_usd": float(args.capital),
        "ending_equity_usd": float(eq.iloc[-1]) if len(eq) else float("nan"),
        "total_pnl_usd": float(cum.iloc[-1]) if len(cum) else float("nan"),
        "portfolio_sharpe_daily": float(sh) if math.isfinite(sh) else None,
        "max_drawdown_frac": dd if math.isfinite(dd) else None,
        "n_closed_trades": int(len(trades_df)),
        "n_nonzero_qty_trades": int((trades_df["qty"] > 0).sum()) if len(trades_df) else 0,
        "sleeve_stats": sleeve_stats,
        "out_daily_csv": _repo_rel(out_d),
        "out_trades_csv": _repo_rel(out_t),
    }
    meta_path = out_d.parent / "literature_low_corr_portfolio_meta.json"
    meta_path.write_text(json.dumps(meta, indent=2), encoding="utf-8")

    print(json.dumps(meta, indent=2), flush=True)
    print(f"\nWrote {out_d}\nWrote {out_t}\nWrote {meta_path}", flush=True)


if __name__ == "__main__":
    main()
