"""
Canonical **Theta + margin + daily MTM** portfolio evaluation (agent-facing API).

Use this instead of summing realized-only daily PnL (``combine_lit_stack_sleeves``,
``portfolio_top_strategies``) when reporting **return, Sharpe, or max drawdown** for
multi-sleeve option books.

Primary risk metric: ``max_drawdown_frac_mtm`` on ``equity_mtm_usd``.
Secondary (reference only): ``max_drawdown_frac_realized`` on exit-day realized equity.
"""
from __future__ import annotations

import json
import math
import time
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Literal

import pandas as pd

from RenTech.core.options_data_loader import OptionChain
from RenTech.core.theta_chunks_loader import ThetaChunksLoader, theta_chunks_date_bounds
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.margin_portfolio_simulator import (
    MarginPortfolioConfig,
    MarginSleeveSpec,
    VrpReplayTrade,
    map_vrp_closed_trades,
    run_margin_portfolio,
)
from RenTech.strategy_stack.diverse_theta_strategies_v1.panel import augment_research_panel
from RenTech.strategy_stack.diverse_theta_strategies_v1.runner import (
    _load_strategy_module,
    _trade_fn,
)
from RenTech.strategy_stack.literature_search_agent import _compile_signal, _compile_trade
from RenTech.strategy_stack.literature_strategy_catalog import StrategySpec, build_catalog_100
from RenTech.strategy_stack.vrp_backtester import VRPBacktester
from RenTech.strategy_stack.vrp_strategy_config import (
    DEFAULT_STRATEGY_CONFIG_PATH,
    apply_strategy_params_to_vrp_backtester_module,
    load_strategy_config_file,
)

PresetName = Literal["lit4", "lit4-vrp", "top9-d", "best-ideas-spy"]

PRESET_LIT4 = ("S055", "S057", "S059", "S089")
PRESET_TOP9_D = ("D039", "D018", "D095", "D081", "D057", "D065", "D041", "D046", "D022")
# Best Ideas literature + curated diverse skew/put sleeves (see BEST_IDEAS.md).
PRESET_D6_BEST = ("D095", "D081", "D039", "D018", "D041", "D065")
PRESET_BEST_IDEAS_SPY = PRESET_LIT4 + PRESET_D6_BEST

PRESET_SIDS: dict[str, tuple[str, ...]] = {
    "lit4": PRESET_LIT4,
    "lit4-vrp": PRESET_LIT4,
    "top9-d": PRESET_TOP9_D,
    "best-ideas-spy": PRESET_BEST_IDEAS_SPY,
}
PRESET_WITH_VRP: frozenset[str] = frozenset({"lit4-vrp", "best-ideas-spy"})


def json_safe(x: Any) -> Any:
    if isinstance(x, dict):
        return {str(k): json_safe(v) for k, v in x.items()}
    if isinstance(x, list):
        return [json_safe(v) for v in x]
    if isinstance(x, float):
        return x if math.isfinite(x) else None
    return x


def resolve_end_date(end: str, theta_dir: Path) -> str:
    e = str(end).strip()
    if e:
        return e
    _, d1 = theta_chunks_date_bounds(theta_dir.expanduser().resolve())
    return d1.strftime("%Y-%m-%d")


def resolve_sids(
    *,
    preset: str | None,
    sids: tuple[str, ...] | list[str] | None,
) -> tuple[tuple[str, ...], bool]:
    """Return (sids, with_vrp)."""
    if preset:
        key = str(preset).strip().lower()
        if key not in PRESET_SIDS:
            raise ValueError(f"Unknown preset {preset!r}; choose from {sorted(PRESET_SIDS)}")
        return PRESET_SIDS[key], key in PRESET_WITH_VRP
    if not sids:
        raise ValueError("Provide --preset or at least one --sid")
    normalized = tuple(str(x).strip().upper() for x in sids)
    return normalized, False


def partition_sids(sids: tuple[str, ...]) -> tuple[tuple[str, ...], tuple[str, ...]]:
    lit: list[str] = []
    diverse: list[str] = []
    catalog = {s.sid for s in build_catalog_100()}
    for sid in sids:
        if sid.startswith("S"):
            if sid not in catalog:
                raise ValueError(f"Unknown literature sid {sid!r}")
            lit.append(sid)
        elif sid.startswith("D") and len(sid) == 4 and sid[1:].isdigit():
            diverse.append(sid)
        else:
            raise ValueError(f"Sid {sid!r}: use S### (literature catalog) or D### (diverse v1)")
    return tuple(lit), tuple(diverse)


def _diverse_mod_index(sid: str) -> int:
    s = sid.strip().upper()
    if not s.startswith("D") or len(s) != 4:
        raise ValueError(f"Bad diverse sid {sid!r}")
    return int(s[1:])


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


def build_literature_margin_sleeves(
    lit_sids: tuple[str, ...],
    *,
    panel: pd.DataFrame,
    iv_atm: dict,
    skew: dict,
    n_contracts: list[int],
    spy_wide: pd.DataFrame,
    stacking: bool = False,
    max_concurrent: int | None = None,
) -> list[MarginSleeveSpec]:
    catalog = _catalog_by_sid()
    out: list[MarginSleeveSpec] = []
    for sid in lit_sids:
        spec = catalog[sid]
        sig_fn = _compile_signal(spec, panel, iv_atm, skew, n_contracts, spy_wide)
        tfn = _compile_trade(spec)
        hold = int(spec.hold)
        if max_concurrent is not None:
            mc = max(1, int(max_concurrent))
        elif stacking:
            mc = max(1, hold)
        else:
            mc = 1

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

        out.append(
            MarginSleeveSpec(
                sid=str(spec.sid),
                trade_kind=str(spec.trade_kind),
                trade_params=tuple(spec.trade_params),
                hold_sessions=hold,
                signal=_sig,
                trade_fn=tfn,
                title=str(spec.description),
                max_concurrent=mc,
            )
        )
    return out


def build_diverse_margin_sleeves(
    diverse_sids: tuple[str, ...],
    ctx: ResearchContext,
    *,
    stacking: bool = False,
    max_concurrent: int | None = None,
) -> list[MarginSleeveSpec]:
    out: list[MarginSleeveSpec] = []
    for sid in diverse_sids:
        k = _diverse_mod_index(sid)
        mod = _load_strategy_module(k)
        meta = dict(getattr(mod, "META"))
        wants = getattr(mod, "wants_entry")
        tk = str(getattr(mod, "TRADE_KIND"))
        tp = tuple(getattr(mod, "TRADE_PARAMS"))
        hold = int(getattr(mod, "HOLD_SESSIONS"))
        tfn = _trade_fn(tk, tp)
        if max_concurrent is not None:
            mc = max(1, int(max_concurrent))
        elif stacking:
            mc = max(1, hold)
        else:
            mc = 1

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

        out.append(
            MarginSleeveSpec(
                sid=str(meta.get("sid", sid)),
                trade_kind=tk,
                trade_params=tp,
                hold_sessions=hold,
                signal=_sig,
                trade_fn=tfn,
                title=str(meta.get("title", "")),
                max_concurrent=mc,
            )
        )
    return out


def run_vrp_backtest_for_eval(
    theta_dir: Path,
    days: list[pd.Timestamp],
    spy_wide: pd.DataFrame,
    capital: float,
    *,
    low_dd_overlap: bool = True,
    overlap_slices: int = 1,
    disable_pmcc: bool = False,
    r2_vix_scale_contracts: bool = False,
    r2_term_structure_gate: bool = False,
    show_progress: bool = True,
) -> list[Any]:
    ld = ThetaChunksLoader(theta_dir, spy_df=spy_wide)
    cfg = load_strategy_config_file(DEFAULT_STRATEGY_CONFIG_PATH)
    apply_strategy_params_to_vrp_backtester_module(cfg.strategy_params)
    bt = VRPBacktester(
        ld,
        initial_capital=float(capital),
        spy_df=spy_wide,
        vol_risk_scaling=False,
        r2_crossover_filters=bool(low_dd_overlap),
        dd_risk_scaling=bool(low_dd_overlap),
        sleeve_risk_fractions=dict(cfg.sleeve_risk_fractions),
        overlap_portfolio=True,
        overlap_slice_contracts=max(1, int(overlap_slices)),
        disable_pmcc=bool(disable_pmcc),
        r2_vix_scale_contracts=bool(r2_vix_scale_contracts),
        r2_term_structure_gate=bool(r2_term_structure_gate),
    )
    if show_progress:
        vrp_tags = []
        if disable_pmcc:
            vrp_tags.append("no_pmcc")
        if r2_vix_scale_contracts:
            vrp_tags.append("vix_scale")
        if r2_term_structure_gate:
            vrp_tags.append("term_gate")
        tag_s = f" flags={','.join(vrp_tags)}" if vrp_tags else ""
        print(
            f"  VRP backtest ({len(days)} sessions, low_dd_overlap={low_dd_overlap}, "
            f"slices={overlap_slices}{tag_s}) …",
            flush=True,
        )
    bt.run_backtest(trading_days=days, show_progress=show_progress)
    return list(bt.trade_log)


def enrich_meta_with_realized_metrics(
    daily_df: pd.DataFrame,
    meta: dict[str, Any],
    capital_start_usd: float,
) -> dict[str, Any]:
    """Add realized-only equity stats (reference; MTM remains primary)."""
    cap = float(capital_start_usd)
    if not len(daily_df):
        return meta
    real_eq = cap + daily_df["cumulative_realized_usd"].astype(float)
    dd_real = float((real_eq / real_eq.cummax() - 1.0).min())
    meta = dict(meta)
    meta["ending_equity_realized_usd"] = float(real_eq.iloc[-1])
    meta["total_return_pct_realized"] = float(real_eq.iloc[-1] / cap - 1.0) * 100.0
    meta["max_drawdown_frac_realized"] = dd_real if math.isfinite(dd_real) else None
    meta["max_drawdown_pct_realized"] = (
        float(dd_real * 100.0) if math.isfinite(dd_real) else None
    )
    meta["evaluation_engine"] = "theta_margin_mtm_v1"
    meta["primary_risk_metric"] = "max_drawdown_frac_mtm"
    meta["primary_return_metric"] = "ending_equity_mtm_usd"
    return meta


@dataclass
class ThetaMarginEvalConfig:
    """Inputs for :func:`run_theta_margin_evaluation`."""

    capital_usd: float = 100_000.0
    start: str = "2016-01-04"
    end: str = ""
    max_days: int = 0
    theta_dir: Path = field(default_factory=lambda: L._DEFAULT_THETA)
    literature_sids: tuple[str, ...] = ()
    diverse_sids: tuple[str, ...] = ()
    with_vrp: bool = False
    vrp_low_dd_overlap: bool = True
    vrp_overlap_slices: int = 1
    vrp_disable_pmcc: bool = False
    vrp_r2_vix_scale_contracts: bool = False
    vrp_r2_term_structure_gate: bool = False
    regt_short_put_mult: float = 1.0
    max_margin_utilization: float = 1.0
    qty_per_trade: int = 1
    stacking: bool = False
    max_concurrent_per_sid: int | None = None
    sleeve_entry_order: tuple[str, ...] | None = None
    show_vrp_progress: bool = True


@dataclass
class ThetaMarginEvalResult:
    daily: pd.DataFrame
    trades: list[dict[str, Any]]
    meta: dict[str, Any]
    days: list[pd.Timestamp]
    wall_s: float


def run_theta_margin_evaluation(cfg: ThetaMarginEvalConfig) -> ThetaMarginEvalResult:
    """
    Run overlapping margin + daily MTM simulation for literature and/or diverse sleeves.

  Optional VRP sleeve runs ``VRPBacktester`` then replays closed trades with intraday MTM.
    """
    t0 = time.perf_counter()
    theta_dir = Path(cfg.theta_dir).expanduser().resolve()
    end = resolve_end_date(cfg.end, theta_dir)

    days, panel0, get_chain, iv_atm, skew, n_contracts, spy_wide = L.prepare_theta_research_context(
        theta_dir=theta_dir,
        capital=float(cfg.capital_usd),
        start=str(cfg.start).strip(),
        end=end,
        max_days=int(cfg.max_days),
    )
    panel = augment_research_panel(panel0)
    idx = pd.DatetimeIndex([L._norm(d) for d in days])

    specs: list[MarginSleeveSpec] = []
    if cfg.literature_sids:
        specs.extend(
            build_literature_margin_sleeves(
                cfg.literature_sids,
                panel=panel,
                iv_atm=iv_atm,
                skew=skew,
                n_contracts=n_contracts,
                spy_wide=spy_wide,
                stacking=bool(cfg.stacking),
                max_concurrent=cfg.max_concurrent_per_sid,
            )
        )
    if cfg.diverse_sids:
        ctx = ResearchContext(days, panel, get_chain, iv_atm, skew, n_contracts)
        specs.extend(
            build_diverse_margin_sleeves(
                cfg.diverse_sids,
                ctx,
                stacking=bool(cfg.stacking),
                max_concurrent=cfg.max_concurrent_per_sid,
            )
        )

    vrp_mapped: list[VrpReplayTrade] = []
    vrp_source_n = 0
    vrp_skipped = 0
    if cfg.with_vrp:
        vrp_log = run_vrp_backtest_for_eval(
            theta_dir,
            days,
            spy_wide,
            float(cfg.capital_usd),
            low_dd_overlap=bool(cfg.vrp_low_dd_overlap),
            overlap_slices=int(cfg.vrp_overlap_slices),
            disable_pmcc=bool(cfg.vrp_disable_pmcc),
            r2_vix_scale_contracts=bool(cfg.vrp_r2_vix_scale_contracts),
            r2_term_structure_gate=bool(cfg.vrp_r2_term_structure_gate),
            show_progress=bool(cfg.show_vrp_progress),
        )
        vrp_source_n = len(vrp_log)
        vrp_mapped, vrp_skipped = map_vrp_closed_trades(vrp_log, idx)
        if cfg.show_vrp_progress:
            print(
                f"  VRP trades mapped: {len(vrp_mapped)} "
                f"(skipped {vrp_skipped} outside session grid)",
                flush=True,
            )

    sid_order = [s.sid for s in specs]
    if cfg.with_vrp:
        sid_order = list(sid_order) + ["VRP"]
    order = tuple(cfg.sleeve_entry_order) if cfg.sleeve_entry_order else tuple(sid_order)

    sim_cfg = MarginPortfolioConfig(
        capital_start_usd=float(cfg.capital_usd),
        regt_short_put_mult=float(cfg.regt_short_put_mult),
        max_margin_utilization=float(cfg.max_margin_utilization),
        sleeve_entry_order=order,
        qty_per_trade=max(1, int(cfg.qty_per_trade)),
        vrp_sid="VRP",
    )

    daily_df, trades, meta = run_margin_portfolio(
        days,
        panel,
        get_chain,
        specs,
        sim_cfg,
        vrp_trades=vrp_mapped if cfg.with_vrp else None,
    )

    meta.update(
        {
            "literature_sids": list(cfg.literature_sids),
            "diverse_sids": list(cfg.diverse_sids),
            "with_vrp": bool(cfg.with_vrp),
            "stacking": bool(cfg.stacking),
            "max_concurrent_per_sid": cfg.max_concurrent_per_sid,
            "vrp_low_dd_overlap": bool(cfg.vrp_low_dd_overlap),
            "vrp_overlap_slices": int(cfg.vrp_overlap_slices),
            "vrp_disable_pmcc": bool(cfg.vrp_disable_pmcc),
            "vrp_r2_vix_scale_contracts": bool(cfg.vrp_r2_vix_scale_contracts),
            "vrp_r2_term_structure_gate": bool(cfg.vrp_r2_term_structure_gate),
            "vrp_closed_trades_source": int(vrp_source_n),
            "vrp_trades_mapped": len(vrp_mapped),
            "vrp_trades_skipped": int(vrp_skipped),
            "theta_dir": str(theta_dir),
            "first_day": str(days[0].date()) if days else "",
            "last_day": str(days[-1].date()) if days else "",
            "command_hint": (
                "python -m RenTech.strategy_stack.diverse_theta_strategies_v1.evaluate_theta_margin"
            ),
        }
    )
    meta = enrich_meta_with_realized_metrics(daily_df, meta, float(cfg.capital_usd))
    wall = time.perf_counter() - t0
    meta["theta_precompute_and_sim_wall_s"] = round(wall, 3)

    return ThetaMarginEvalResult(
        daily=daily_df,
        trades=trades,
        meta=meta,
        days=days,
        wall_s=wall,
    )


def write_theta_margin_eval_outputs(
    result: ThetaMarginEvalResult,
    *,
    out_daily: Path,
    out_trades: Path,
    out_meta: Path,
) -> tuple[Path, Path, Path]:
    out_daily = Path(out_daily).expanduser().resolve()
    out_trades = Path(out_trades).expanduser().resolve()
    out_meta = Path(out_meta).expanduser().resolve()
    out_daily.parent.mkdir(parents=True, exist_ok=True)
    result.daily.to_csv(out_daily, index=False)
    pd.DataFrame(result.trades).to_csv(out_trades, index=False)
    out_meta.write_text(json.dumps(json_safe(result.meta), indent=2), encoding="utf-8")
    return out_daily, out_trades, out_meta


def default_output_paths(
    prefix: Path,
    *,
    slug: str = "eval",
) -> tuple[Path, Path, Path]:
    """``prefix`` may be a directory or a path stem without extension."""
    p = Path(prefix).expanduser()
    if p.suffix.lower() in (".csv", ".json"):
        stem = p.with_suffix("")
        return (
            Path(f"{stem}_daily.csv"),
            Path(f"{stem}_trades.csv"),
            Path(f"{stem}_meta.json"),
        )
    if str(p).endswith("_"):
        base = str(p)
    else:
        base = str(p / slug) if p.is_dir() or not p.name or "." not in p.name else str(p)
    return (
        Path(f"{base}_daily.csv"),
        Path(f"{base}_trades.csv"),
        Path(f"{base}_meta.json"),
    )


def config_from_cli_args(
    *,
    preset: str | None,
    sids: list[str],
    with_vrp: bool,
    capital: float,
    start: str,
    end: str,
    max_days: int,
    theta_dir: Path,
    stacking: bool,
    max_concurrent: int | None,
    vrp_low_dd_overlap: bool,
    vrp_overlap_slices: int,
    vrp_disable_pmcc: bool = False,
    vrp_r2_vix_scale_contracts: bool = False,
    vrp_r2_term_structure_gate: bool = False,
    regt_short_put_mult: float,
    max_margin_utilization: float,
    qty_per_trade: int,
    show_vrp_progress: bool,
) -> ThetaMarginEvalConfig:
    resolved, preset_vrp = resolve_sids(preset=preset, sids=tuple(sids) if sids else None)
    use_vrp = bool(with_vrp) or preset_vrp
    lit, diverse = partition_sids(resolved)
    return ThetaMarginEvalConfig(
        capital_usd=float(capital),
        start=str(start),
        end=str(end),
        max_days=int(max_days),
        theta_dir=Path(theta_dir),
        literature_sids=lit,
        diverse_sids=diverse,
        with_vrp=use_vrp,
        vrp_low_dd_overlap=bool(vrp_low_dd_overlap),
        vrp_overlap_slices=int(vrp_overlap_slices),
        vrp_disable_pmcc=bool(vrp_disable_pmcc),
        vrp_r2_vix_scale_contracts=bool(vrp_r2_vix_scale_contracts),
        vrp_r2_term_structure_gate=bool(vrp_r2_term_structure_gate),
        regt_short_put_mult=float(regt_short_put_mult),
        max_margin_utilization=float(max_margin_utilization),
        qty_per_trade=int(qty_per_trade),
        stacking=bool(stacking),
        max_concurrent_per_sid=max_concurrent,
        show_vrp_progress=bool(show_vrp_progress),
    )
