"""
Shared **VRP strategy + sleeve risk** configuration for parity between:

* ``vrp_backtester.py`` (Theta / iVol / synthetic backtests)
* ``live_ibkr_trader.py`` (IBKR)

Single JSON file (default ``sleeve_risk_fractions.json``) holds:

* ``sleeve_risk_fractions`` — per-regime risk fractions (same keys as ``Regime``)
* ``overlay_risk_fractions`` — reserved for future in-engine overlays (passed to ``VRPBacktester``)
* ``overlay_risk_cap_frac`` / ``total_risk_cap_frac`` — optional caps
* ``strategy_params`` — optional scalars / leg-spec lists mirrored onto ``vrp_backtester`` module globals
"""

from __future__ import annotations

import importlib
import json
import math
import types
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Literal

_REPO = Path(__file__).resolve().parents[2]
DEFAULT_STRATEGY_CONFIG_PATH = _REPO / "RenTech" / "strategy_stack" / "sleeve_risk_fractions.json"

RegimeStr = Literal["pmcc", "diagonal", "r2_spread", "naked", "credit_spread"]
_VALID_SLEEVES: set[str] = {"pmcc", "diagonal", "r2_spread", "naked", "credit_spread"}


def _validate_sleeve_fracs(raw: dict[str, Any]) -> dict[str, float]:
    out: dict[str, float] = {}
    for k, v in raw.items():
        kk = str(k).strip()
        if kk not in _VALID_SLEEVES:
            raise ValueError(f"Unknown sleeve '{kk}'; valid: {sorted(_VALID_SLEEVES)}")
        fv = float(v)
        if not (math.isfinite(fv) and 0.0 < fv <= 1.0):
            raise ValueError(f"sleeve_risk_fractions['{kk}'] must be in (0, 1], got {v}")
        out[kk] = fv
    return out


def _validate_overlay_fracs(raw: dict[str, Any]) -> dict[str, float]:
    out: dict[str, float] = {}
    for k, v in raw.items():
        kk = str(k).strip()
        if not kk:
            raise ValueError("overlay_risk_fractions key cannot be empty")
        fv = float(v)
        if not (math.isfinite(fv) and 0.0 < fv <= 1.0):
            raise ValueError(f"overlay_risk_fractions['{kk}'] must be in (0, 1], got {v}")
        out[kk] = fv
    return out


def _coerce_credit_leg_specs(raw: Any, *, label: str) -> list[tuple[int, float, str]]:
    if not isinstance(raw, list) or len(raw) < 2:
        raise ValueError(f"{label} must be a list of at least two [dte, delta, action] rows")
    out: list[tuple[int, float, str]] = []
    for i, row in enumerate(raw):
        if not isinstance(row, (list, tuple)) or len(row) < 3:
            raise ValueError(f"{label}[{i}] must be [dte, delta, action]")
        dte = int(row[0])
        delta = float(row[1])
        act = str(row[2]).strip().lower()
        if act not in ("buy", "sell"):
            raise ValueError(f"{label}[{i}] action must be buy/sell, got {row[2]!r}")
        out.append((dte, delta, act))
    return out


@dataclass(frozen=True)
class VRPStrategyConfigFile:
    """Parsed ``sleeve_risk_fractions.json`` (and friends)."""

    path: Path
    sleeve_risk_fractions: dict[str, float]
    overlay_risk_fractions: dict[str, float]
    overlay_risk_cap_frac: float | None
    total_risk_cap_frac: float | None
    strategy_params: dict[str, Any]


def load_strategy_config_file(path: str | Path | None = None) -> VRPStrategyConfigFile:
    p = Path(path).expanduser() if path is not None else DEFAULT_STRATEGY_CONFIG_PATH
    if not p.is_file():
        raise FileNotFoundError(f"Strategy config not found: {p}")
    payload = json.loads(p.read_text(encoding="utf-8"))
    if not isinstance(payload, dict):
        raise ValueError(f"Config must be a JSON object: {p}")

    sleeve_raw = payload.get("sleeve_risk_fractions", {})
    if not isinstance(sleeve_raw, dict):
        raise ValueError("'sleeve_risk_fractions' must be an object")
    sleeves = _validate_sleeve_fracs(sleeve_raw)

    ov_raw = payload.get("overlay_risk_fractions", {})
    if ov_raw is not None and not isinstance(ov_raw, dict):
        raise ValueError("'overlay_risk_fractions' must be an object or omitted")
    overlays = _validate_overlay_fracs(ov_raw or {})

    oc = payload.get("overlay_risk_cap_frac")
    overlay_cap = float(oc) if oc is not None else None
    if overlay_cap is not None and not (0.0 < overlay_cap <= 1.0):
        raise ValueError("overlay_risk_cap_frac must be in (0, 1]")

    tc = payload.get("total_risk_cap_frac")
    total_cap = float(tc) if tc is not None else None
    if total_cap is not None and not (0.0 < total_cap <= 1.0):
        raise ValueError("total_risk_cap_frac must be in (0, 1]")

    sp = payload.get("strategy_params", {})
    if sp is not None and not isinstance(sp, dict):
        raise ValueError("'strategy_params' must be an object or omitted")
    strategy_params = dict(sp or {})

    return VRPStrategyConfigFile(
        path=p.resolve(),
        sleeve_risk_fractions=sleeves,
        overlay_risk_fractions=overlays,
        overlay_risk_cap_frac=overlay_cap,
        total_risk_cap_frac=total_cap,
        strategy_params=strategy_params,
    )


def load_sleeve_risk_fractions_from_json(path: str | Path) -> dict[str, float]:
    """Backward-compatible helper: return only ``sleeve_risk_fractions`` mapping."""
    cfg = load_strategy_config_file(path)
    return dict(cfg.sleeve_risk_fractions)


def apply_strategy_params_to_vrp_backtester_module(params: dict[str, Any] | None) -> None:
    """
    Mutate ``vrp_backtester`` module-level constants from ``strategy_params`` JSON.

    This keeps ``VRPBacktester`` logic unchanged while allowing one JSON file to drive
    both Theta backtests and live trading parity.
    """
    if not params:
        return
    import RenTech.strategy_stack.vrp_backtester as m

    simple_float_int = {
        "vix_r1_max": float,
        "vix_r2_max": float,
        "vix_r3_max": float,
        "r1_strangle_dte": int,
        "r1_call_delta": float,
        "r1_put_delta": float,
        "r1_time_stop_days": int,
        "r1_tp_frac": float,
        "r1_sl_frac": float,
        "risk_fraction_per_trade": float,
        "r2_put_spread_risk_fraction": float,
        "r2_risk_per_contract": float,
        "r2_tp_per_contract": float,
        "r2_sl_per_contract": float,
        "r2_time_stop_days": int,
        "r2_spread_tp_frac": float,
        "r2_spread_time_stop_days": int,
        "r3_tp_frac": float,
        "r3_sl_frac": float,
        "r3_time_stop_days": int,
        "r4_tp_frac": float,
        "r4_time_stop_days": int,
        "sizing_min_max_risk_usd": float,
        "contract_multiplier": float,
    }

    for json_key, attr in [
        ("vix_r1_max", "VIX_R1_MAX"),
        ("vix_r2_max", "VIX_R2_MAX"),
        ("vix_r3_max", "VIX_R3_MAX"),
        ("r1_strangle_dte", "R1_STRANGLE_DTE"),
        ("r1_call_delta", "R1_CALL_DELTA"),
        ("r1_put_delta", "R1_PUT_DELTA"),
        ("r1_time_stop_days", "R1_TIME_STOP_DAYS"),
        ("r1_tp_frac", "R1_TP_FRAC"),
        ("r1_sl_frac", "R1_SL_FRAC"),
        ("risk_fraction_per_trade", "RISK_FRACTION_PER_TRADE"),
        ("r2_put_spread_risk_fraction", "R2_PUT_SPREAD_RISK_FRACTION"),
        ("r2_risk_per_contract", "R2_RISK_PER_CONTRACT"),
        ("r2_tp_per_contract", "R2_TP_PER_CONTRACT"),
        ("r2_sl_per_contract", "R2_SL_PER_CONTRACT"),
        ("r2_time_stop_days", "R2_TIME_STOP_DAYS"),
        ("r2_spread_tp_frac", "R2_SPREAD_TP_FRAC"),
        ("r2_spread_time_stop_days", "R2_SPREAD_TIME_STOP_DAYS"),
        ("r3_tp_frac", "R3_TP_FRAC"),
        ("r3_sl_frac", "R3_SL_FRAC"),
        ("r3_time_stop_days", "R3_TIME_STOP_DAYS"),
        ("r4_tp_frac", "R4_TP_FRAC"),
        ("r4_time_stop_days", "R4_TIME_STOP_DAYS"),
        ("sizing_min_max_risk_usd", "SIZING_MIN_MAX_RISK_USD"),
        ("contract_multiplier", "CONTRACT_MULTIPLIER"),
    ]:
        if json_key not in params:
            continue
        caster = simple_float_int[json_key]
        setattr(m, attr, caster(params[json_key]))

    # R2 diagonal two-leg specs: either full list or short/long scalars
    if "r2_diag_leg_specs" in params:
        m.R2_DIAG_LEG_SPECS = _coerce_credit_leg_specs(params["r2_diag_leg_specs"], label="r2_diag_leg_specs")
    elif all(k in params for k in ("r2_short_dte", "r2_short_delta", "r2_long_dte", "r2_long_delta")):
        m.R2_DIAG_LEG_SPECS = [
            (int(params["r2_short_dte"]), float(params["r2_short_delta"]), "sell"),
            (int(params["r2_long_dte"]), float(params["r2_long_delta"]), "buy"),
        ]

    if "r2_credit_leg_specs" in params:
        m.R2_CREDIT_LEG_SPECS = _coerce_credit_leg_specs(
            params["r2_credit_leg_specs"], label="r2_credit_leg_specs"
        )

    if "r2_spread_target_dte" in params:
        m.R2_SPREAD_TARGET_DTE = int(params["r2_spread_target_dte"])
    if "r2_spread_short_delta" in params:
        m.R2_SPREAD_SHORT_DELTA = float(params["r2_spread_short_delta"])
    if "r2_spread_width_frac_of_spot" in params:
        m.R2_SPREAD_WIDTH_FRAC_OF_SPOT = float(params["r2_spread_width_frac_of_spot"])

    if "r3_credit_leg_specs" in params:
        m.R3_CREDIT_LEG_SPECS = _coerce_credit_leg_specs(
            params["r3_credit_leg_specs"], label="r3_credit_leg_specs"
        )

    if "r4_credit_leg_specs" in params:
        m.R4_CREDIT_LEG_SPECS = _coerce_credit_leg_specs(
            params["r4_credit_leg_specs"], label="r4_credit_leg_specs"
        )


def sync_live_constants_from_vrp_backtester_module(target: types.ModuleType | None = None) -> None:
    """
    Copy mirrored module-level constants from ``vrp_backtester`` into the live trader module globals.

    Call **after** :func:`apply_strategy_params_to_vrp_backtester_module`.

    Parameters
    ----------
    target
        Module object to patch (use ``sys.modules[__name__]`` when running ``live_ibkr_trader.py``
        as a script so ``__main__`` receives the constants).
    """
    import RenTech.strategy_stack.vrp_backtester as vb

    live = target or importlib.import_module("RenTech.strategy_stack.live_ibkr_trader")
    g = live.__dict__
    names = [
        "CONTRACT_MULTIPLIER",
        "SIZING_MIN_MAX_RISK_USD",
        "VIX_R1_MAX",
        "VIX_R2_MAX",
        "VIX_R3_MAX",
        "R1_TIME_STOP_DAYS",
        "R1_TP_FRAC",
        "R1_SL_FRAC",
        "R1_STRANGLE_DTE",
        "R1_CALL_DELTA",
        "R1_PUT_DELTA",
        "R2_RISK_PER_CONTRACT",
        "R2_TP_PER_CONTRACT",
        "R2_SL_PER_CONTRACT",
        "R2_TIME_STOP_DAYS",
        "R3_TP_FRAC",
        "R3_SL_FRAC",
        "R3_TIME_STOP_DAYS",
        "R4_TP_FRAC",
        "R4_TIME_STOP_DAYS",
        "RISK_FRACTION_PER_TRADE",
    ]
    for n in names:
        if hasattr(vb, n):
            g[n] = getattr(vb, n)

    # R2 diagonal DTE/delta live aliases track R2_DIAG_LEG_SPECS
    if len(vb.R2_DIAG_LEG_SPECS) >= 2:
        g["R2_SHORT_DTE"] = int(vb.R2_DIAG_LEG_SPECS[0][0])
        g["R2_SHORT_DELTA"] = float(vb.R2_DIAG_LEG_SPECS[0][1])
        g["R2_LONG_DTE"] = int(vb.R2_DIAG_LEG_SPECS[1][0])
        g["R2_LONG_DELTA"] = float(vb.R2_DIAG_LEG_SPECS[1][1])

    # R3/R4 short legs: first spec row drives live single-expiry picker
    if len(vb.R3_CREDIT_LEG_SPECS) >= 1:
        g["R3_TARGET_DTE"] = int(vb.R3_CREDIT_LEG_SPECS[0][0])
        g["R3_TARGET_DELTA"] = float(vb.R3_CREDIT_LEG_SPECS[0][1])
    if len(vb.R4_CREDIT_LEG_SPECS) >= 2:
        g["R4_SHORT_DTE"] = int(vb.R4_CREDIT_LEG_SPECS[0][0])
        g["R4_SHORT_DELTA"] = float(vb.R4_CREDIT_LEG_SPECS[0][1])
        g["R4_LONG_DELTA"] = float(vb.R4_CREDIT_LEG_SPECS[1][1])

    if hasattr(vb, "R2_SPREAD_TARGET_DTE"):
        g["R2B_SHORT_DTE"] = int(vb.R2_SPREAD_TARGET_DTE)
    if hasattr(vb, "R2_SPREAD_SHORT_DELTA"):
        g["R2B_SHORT_DELTA"] = float(vb.R2_SPREAD_SHORT_DELTA)
    if hasattr(vb, "R2_SPREAD_WIDTH_FRAC_OF_SPOT"):
        g["R2B_WIDTH_FRAC_OF_SPOT"] = float(vb.R2_SPREAD_WIDTH_FRAC_OF_SPOT)
