"""
Broker-style **multi-sleeve** simulation for Theta SPY option research + VRP trade replay.

- **Overlapping margin:** sum ``margin_reserved_usd * qty`` across all open positions.
- **BP gate:** ``margin_used + margin_new <= equity_mtm * max_margin_utilization``.
- **Daily MTM:** realized + mark-to-market on open legs (conservative bid/ask).
- **Literature sleeves:** at most ``max_concurrent`` open trades per ``sid`` (default 1).
- **VRP:** replay ``ClosedTrade`` rows from a full ``VRPBacktester`` run (many concurrent
  positions allowed on sid ``VRP``).
"""
from __future__ import annotations

import json
import math
from dataclasses import dataclass, field
from typing import Any, Callable, Iterable

import numpy as np
import pandas as pd

from RenTech.core.options_backtest import find_contract_in_chain
from RenTech.core.options_data_loader import OptionChain, OptionContract
from RenTech.strategy_stack import research_literature_theta_strategies as L
from RenTech.strategy_stack.literature_low_corr_portfolio import _margin_per_contract

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


def _parse_expiry(s: str) -> pd.Timestamp:
    return pd.Timestamp(str(s).strip()).normalize()


def _intrinsic_per_share(right: str, strike: float, spy: float) -> float:
    r = str(right).upper()[:1]
    k = float(strike)
    s = float(spy)
    if r == "C":
        return max(0.0, s - k)
    return max(0.0, k - s)


def _mark_leg_mtm_usd(leg: dict[str, Any], chain: OptionChain, spy: float, as_of: pd.Timestamp) -> float:
    mult = float(L.MULT)
    exp = _parse_expiry(str(leg["expiry"]))
    strike = float(leg["strike"])
    right = str(leg["right"]).upper()[:1]
    ot = "C" if right == "C" else "P"
    pos = str(leg["position"]).lower()
    day = L._norm(as_of)
    expired = day >= exp
    live = None if expired else find_contract_in_chain(chain, exp, strike, ot)

    if pos == "short":
        entry_px = float(leg["entry_bid"])
        if expired:
            exit_px = float(leg.get("exit_intrinsic_per_sh") or _intrinsic_per_share(right, strike, spy))
        elif live is not None:
            exit_px = float(live.ask)
        else:
            exit_px = _intrinsic_per_share(right, strike, spy)
        return float(entry_px - exit_px) * mult

    entry_px = float(leg["entry_ask"])
    if expired:
        mark = float(leg.get("exit_intrinsic_per_sh") or _intrinsic_per_share(right, strike, spy))
    elif live is not None:
        mark = float(live.bid)
    else:
        mark = _intrinsic_per_share(right, strike, spy)
    return float(mark - entry_px) * mult


def _position_mtm_usd(legs: list[dict[str, Any]], chain: OptionChain, spy: float, as_of: pd.Timestamp) -> float:
    if not legs:
        return 0.0
    return float(sum(_mark_leg_mtm_usd(lg, chain, spy, as_of) for lg in legs))


def _leg_row_at_entry(c: OptionContract, position: str, spy: float, as_of: pd.Timestamp) -> dict[str, Any]:
    exp = L._norm(c.expiration)
    if str(position).lower() == "short":
        return {
            "position": "short",
            "right": str(c.option_type).upper()[:1],
            "strike": float(c.strike),
            "expiry": exp.strftime("%Y-%m-%d"),
            "entry_bid": float(c.bid),
            "entry_ask": float(c.ask),
            "exit_bid": None,
            "exit_ask": None,
            "exit_settled": False,
            "exit_intrinsic_per_sh": None,
            "entry_iv": float(c.iv),
            "delta_entry": float(c.delta),
        }
    return {
        "position": "long",
        "right": str(c.option_type).upper()[:1],
        "strike": float(c.strike),
        "expiry": exp.strftime("%Y-%m-%d"),
        "entry_bid": float(c.bid),
        "entry_ask": float(c.ask),
        "exit_bid": None,
        "exit_ask": None,
        "exit_settled": False,
        "exit_intrinsic_per_sh": None,
        "entry_iv": float(c.iv),
        "delta_entry": float(c.delta),
    }


def _vrp_legs_from_chain(
    chain: OptionChain,
    spy: float,
    legs_json: str,
    regime: str,
) -> list[dict[str, Any]]:
    """Build MTM leg rows from VRP ``legs_json`` + same-day chain."""
    try:
        raw = json.loads(legs_json or "[]")
    except json.JSONDecodeError:
        return []
    if not raw:
        return []
    as_of = L._norm(chain.as_of)
    regime_l = str(regime).lower()

    def _find(right: str, strike: float, expiry: str) -> OptionContract | None:
        exp = _parse_expiry(expiry)
        ot = "C" if str(right).upper().startswith("C") else "P"
        return find_contract_in_chain(chain, exp, float(strike), ot)

    if regime_l in ("r2_spread", "credit_spread"):
        puts = [x for x in raw if str(x.get("right", "P")).upper().startswith("P")]
        if len(puts) < 2:
            return []
        puts.sort(key=lambda x: float(x["strike"]), reverse=True)
        short_c = _find("P", float(puts[0]["strike"]), str(puts[0]["expiry"]))
        long_c = _find("P", float(puts[1]["strike"]), str(puts[1]["expiry"]))
        if short_c is None or long_c is None:
            return []
        return [
            _leg_row_at_entry(short_c, "short", spy, as_of),
            _leg_row_at_entry(long_c, "long", spy, as_of),
        ]

    if regime_l == "naked" and raw:
        x = raw[0]
        c = _find(str(x.get("right", "P")), float(x["strike"]), str(x["expiry"]))
        if c is None:
            return []
        return [_leg_row_at_entry(c, "short", spy, as_of)]

    if regime_l == "diagonal" and len(raw) >= 2:
        legs_out: list[dict[str, Any]] = []
        for x in raw:
            c = _find(str(x.get("right", "P")), float(x["strike"]), str(x["expiry"]))
            if c is None:
                continue
            legs_out.append(_leg_row_at_entry(c, "long", spy, as_of))
        return legs_out

    return []


@dataclass
class MarginSleeveSpec:
    sid: str
    trade_kind: str
    trade_params: tuple
    hold_sessions: int
    signal: SignalFn
    trade_fn: TradeFn
    title: str = ""
    max_concurrent: int = 1


@dataclass
class OpenSimTrade:
    sid: str
    entry_idx: int
    exit_idx: int
    trade_kind: str
    trade_params: tuple
    entry_spy: float
    legs: list[dict[str, Any]]
    margin_reserved_usd: float
    trade_fn: TradeFn
    qty: int = 1
    realized_pnl_override: float | None = None
    regime: str = ""
    uid: int = 0


@dataclass
class VrpReplayTrade:
    """Precomputed VRP closed trade mapped to session indices."""

    entry_idx: int
    exit_idx: int
    margin_reserved_usd: float
    realized_pnl_usd: float
    qty: int
    regime: str
    legs_json: str
    exit_reason: str = ""


@dataclass
class MarginPortfolioConfig:
    capital_start_usd: float
    regt_short_put_mult: float = 1.0
    max_margin_utilization: float = 1.0
    sleeve_entry_order: tuple[str, ...] | None = None
    qty_per_trade: int = 1
    vrp_sid: str = "VRP"


def _count_open(opens: list[OpenSimTrade], sid: str) -> int:
    return sum(1 for o in opens if o.sid == sid)


def _try_open(
    opens: list[OpenSimTrade],
    *,
    uid: int,
    sid: str,
    entry_idx: int,
    exit_idx: int,
    trade_kind: str,
    trade_params: tuple,
    entry_spy: float,
    legs: list[dict[str, Any]],
    margin_per_contract: float,
    trade_fn: TradeFn,
    qty: int,
    realized_pnl_override: float | None = None,
    regime: str = "",
) -> OpenSimTrade | None:
    if not legs and realized_pnl_override is None:
        return None
    q = max(1, int(qty))
    return OpenSimTrade(
        sid=sid,
        entry_idx=int(entry_idx),
        exit_idx=int(exit_idx),
        trade_kind=trade_kind,
        trade_params=tuple(trade_params),
        entry_spy=float(entry_spy),
        legs=list(legs),
        margin_reserved_usd=float(margin_per_contract),
        trade_fn=trade_fn,
        qty=q,
        realized_pnl_override=realized_pnl_override,
        regime=str(regime),
        uid=int(uid),
    )


def run_margin_portfolio(
    days: list[pd.Timestamp],
    panel: pd.DataFrame,
    get_chain: Callable[[pd.Timestamp], OptionChain],
    sleeves: Iterable[MarginSleeveSpec],
    cfg: MarginPortfolioConfig,
    *,
    vrp_trades: list[VrpReplayTrade] | None = None,
    vrp_entries_by_day: dict[int, list[VrpReplayTrade]] | None = None,
) -> tuple[pd.DataFrame, list[dict[str, Any]], dict[str, Any]]:
    sleeves_list = list(sleeves)
    by_sid = {s.sid: s for s in sleeves_list}
    order = list(cfg.sleeve_entry_order) if cfg.sleeve_entry_order else [s.sid for s in sleeves_list]

    idx = pd.DatetimeIndex([L._norm(d) for d in days])
    n = len(days)
    day_to_i = {idx[i]: i for i in range(n)}

    if vrp_entries_by_day is None and vrp_trades:
        vrp_entries_by_day = {}
        for vt in vrp_trades:
            vrp_entries_by_day.setdefault(int(vt.entry_idx), []).append(vt)
    vrp_entries_by_day = vrp_entries_by_day or {}

    cumulative_realized = 0.0
    opens: list[OpenSimTrade] = []
    closed_rows: list[dict[str, Any]] = []
    skipped_bp = 0
    skipped_no_legs = 0
    skipped_vrp_no_idx = 0
    uid_ctr = 0

    daily: list[dict[str, Any]] = []

    def _margin_and_unreal(ch: OptionChain, spy: float, d: pd.Timestamp) -> tuple[float, float]:
        mu = 0.0
        ur = 0.0
        for op in opens:
            mu += float(op.margin_reserved_usd) * int(op.qty)
            ur += _position_mtm_usd(op.legs, ch, spy, d) * int(op.qty)
        return mu, ur

    def _bp_allows(margin_new: float, extra_unreal: float) -> bool:
        mu, ur = _margin_and_unreal(ch, spy, d)
        eq = float(cfg.capital_start_usd) + cumulative_realized + ur + extra_unreal
        return (mu + margin_new) <= eq * float(cfg.max_margin_utilization) + 1e-6

    for i in range(n):
        d = days[i]
        row = panel.iloc[i]
        spy = float(row["close"])
        ch = get_chain(d)

        # ----- exits
        still: list[OpenSimTrade] = []
        for op in opens:
            if i < op.exit_idx:
                still.append(op)
                continue
            if op.realized_pnl_override is not None:
                pnl = float(op.realized_pnl_override)
            else:
                ch_e = get_chain(days[op.entry_idx])
                pnl_1x = op.trade_fn(ch_e, ch, op.entry_spy, spy, op.trade_params)
                if pnl_1x is None or not math.isfinite(float(pnl_1x)):
                    continue
                pnl = float(pnl_1x) * int(op.qty)
            cumulative_realized += pnl
            closed_rows.append(
                {
                    "sid": op.sid,
                    "uid": int(op.uid),
                    "entry_date": idx[op.entry_idx].strftime("%Y-%m-%d"),
                    "exit_date": idx[i].strftime("%Y-%m-%d"),
                    "hold_sessions": int(op.exit_idx - op.entry_idx),
                    "qty": int(op.qty),
                    "margin_reserved_usd": float(op.margin_reserved_usd) * int(op.qty),
                    "realized_pnl_usd": float(pnl),
                    "trade_kind": op.trade_kind,
                    "regime": op.regime,
                    "trade_params_json": json.dumps(list(op.trade_params)),
                    "legs_json": json.dumps(op.legs, separators=(",", ":")),
                }
            )
        opens = still

        margin_used, unreal = _margin_and_unreal(ch, spy, d)
        equity_mtm = float(cfg.capital_start_usd) + cumulative_realized + unreal
        bp_util = margin_used / equity_mtm if equity_mtm > 1e-9 else float("nan")

        # ----- VRP scheduled entries (overlap allowed)
        for vt in vrp_entries_by_day.get(i, []):
            legs = _vrp_legs_from_chain(ch, spy, vt.legs_json, vt.regime)
            mpc = float(vt.margin_reserved_usd)
            if mpc <= 0 and legs:
                mpc = float(
                    _margin_per_contract(
                        "vert" if vt.regime in ("r2_spread", "credit_spread") else "put",
                        spy,
                        legs,
                        float(cfg.regt_short_put_mult),
                    )
                )
            qty = max(1, int(vt.qty))
            margin_new = mpc * qty
            extra_ur = _position_mtm_usd(legs, ch, spy, d) * qty if legs else 0.0
            if not _bp_allows(margin_new, extra_ur):
                skipped_bp += 1
                continue
            uid_ctr += 1
            op = _try_open(
                opens,
                uid=uid_ctr,
                sid=str(cfg.vrp_sid),
                entry_idx=i,
                exit_idx=int(vt.exit_idx),
                trade_kind="vrp_replay",
                trade_params=(),
                entry_spy=spy,
                legs=legs,
                margin_per_contract=mpc,
                trade_fn=L._wrap_put((1, -0.2)),
                qty=qty,
                realized_pnl_override=float(vt.realized_pnl_usd),
                regime=str(vt.regime),
            )
            if op is not None:
                opens.append(op)
            else:
                skipped_no_legs += 1

        margin_used, unreal = _margin_and_unreal(ch, spy, d)
        equity_mtm = float(cfg.capital_start_usd) + cumulative_realized + unreal
        bp_util = margin_used / equity_mtm if equity_mtm > 1e-9 else float("nan")

        # ----- signal-driven literature sleeves
        for sid in order:
            if sid not in by_sid:
                continue
            spec = by_sid[sid]
            max_c = int(spec.max_concurrent)
            if max_c > 0 and _count_open(opens, sid) >= max_c:
                continue
            if not ch.contracts:
                continue
            if not spec.signal(i, row, ch, spy):
                continue

            legs = L.collect_trade_legs(spec.trade_kind, spec.trade_params, ch, ch, spy, spy)
            if not legs:
                skipped_no_legs += 1
                continue

            mpc = float(
                _margin_per_contract(
                    str(spec.trade_kind),
                    spy,
                    legs,
                    float(cfg.regt_short_put_mult),
                )
            )
            qty = max(1, int(cfg.qty_per_trade))
            margin_new = mpc * qty
            extra_ur = _position_mtm_usd(legs, ch, spy, d) * qty
            if not _bp_allows(margin_new, extra_ur):
                skipped_bp += 1
                continue

            exit_i = min(n - 1, i + int(spec.hold_sessions))
            uid_ctr += 1
            op = _try_open(
                opens,
                uid=uid_ctr,
                sid=sid,
                entry_idx=i,
                exit_idx=exit_i,
                trade_kind=str(spec.trade_kind),
                trade_params=tuple(spec.trade_params),
                entry_spy=spy,
                legs=legs,
                margin_per_contract=mpc,
                trade_fn=spec.trade_fn,
                qty=qty,
            )
            if op is not None:
                opens.append(op)

        margin_used, unreal = _margin_and_unreal(ch, spy, d)
        equity_mtm = float(cfg.capital_start_usd) + cumulative_realized + unreal
        bp_util = margin_used / equity_mtm if equity_mtm > 1e-9 else float("nan")

        open_sids = sorted({o.sid for o in opens})
        daily.append(
            {
                "date": idx[i].strftime("%Y-%m-%d"),
                "spy_close": spy,
                "equity_mtm_usd": equity_mtm,
                "cumulative_realized_usd": cumulative_realized,
                "unrealized_mtm_usd": unreal,
                "margin_total_usd": margin_used,
                "bp_utilization": bp_util,
                "n_open_positions": len(opens),
                "n_open_vrp": _count_open(opens, str(cfg.vrp_sid)),
                "open_sleeves": ",".join(open_sids),
            }
        )

    df = pd.DataFrame(daily)
    if len(df):
        df["daily_return"] = df["equity_mtm_usd"].pct_change().replace([np.inf, -np.inf], np.nan).fillna(0.0)
        dd = float((df["equity_mtm_usd"] / df["equity_mtm_usd"].cummax() - 1.0).min())
        sh = float(L.sharpe_daily_returns(df["daily_return"])) if len(df) >= 50 else float("nan")
    else:
        dd, sh = float("nan"), float("nan")

    meta = {
        "capital_start_usd": float(cfg.capital_start_usd),
        "ending_equity_mtm_usd": float(df["equity_mtm_usd"].iloc[-1]) if len(df) else float("nan"),
        "total_return_pct": float(df["equity_mtm_usd"].iloc[-1] / cfg.capital_start_usd - 1.0) * 100.0
        if len(df)
        else float("nan"),
        "portfolio_sharpe_daily_mtm": sh if math.isfinite(sh) else None,
        "max_drawdown_frac_mtm": dd if math.isfinite(dd) else None,
        "skipped_entry_insufficient_bp": int(skipped_bp),
        "skipped_entry_no_legs": int(skipped_no_legs),
        "skipped_vrp_no_session_index": int(skipped_vrp_no_idx),
        "regt_short_put_mult": float(cfg.regt_short_put_mult),
        "max_margin_utilization": float(cfg.max_margin_utilization),
        "qty_per_trade": int(cfg.qty_per_trade),
        "n_sessions": int(n),
        "n_closed_trades": int(len(closed_rows)),
        "n_vrp_replay_trades": int(len(vrp_trades) if vrp_trades else 0),
    }
    _ = day_to_i
    return df, closed_rows, meta


def map_vrp_closed_trades(
    closed: list[Any],
    day_index: pd.DatetimeIndex,
) -> tuple[list[VrpReplayTrade], int]:
    """
    Map ``VRPBacktester`` ``ClosedTrade`` rows to session indices. Returns (list, n_skipped).
    """
    day_to_i = {pd.Timestamp(t).normalize(): int(i) for i, t in enumerate(day_index)}
    out: list[VrpReplayTrade] = []
    skipped = 0
    for t in closed:
        ed = pd.Timestamp(t.entry_date).normalize()
        xd = pd.Timestamp(t.exit_date).normalize()
        if ed not in day_to_i or xd not in day_to_i:
            skipped += 1
            continue
        ei = day_to_i[ed]
        xi = day_to_i[xd]
        if xi < ei:
            skipped += 1
            continue
        qty = max(1, int(getattr(t, "qty", 1)))
        mrg = float(t.max_margin) / float(qty) if qty else float(t.max_margin)
        out.append(
            VrpReplayTrade(
                entry_idx=ei,
                exit_idx=xi,
                margin_reserved_usd=float(mrg),
                realized_pnl_usd=float(t.pnl_usd),
                qty=qty,
                regime=str(getattr(t, "regime", "")),
                legs_json=str(getattr(t, "legs_json", "") or ""),
                exit_reason=str(getattr(t, "exit_reason", "")),
            )
        )
    return out, skipped
