"""
Option execution templates for SPX regime state-space actions.

Reuses bid/ask conventions from ``research_literature_theta_strategies``.
"""

from __future__ import annotations

import math
from dataclasses import dataclass
from typing import Callable

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.spx_regime_state_space.classifier import ActionId

MULT = L.MULT
_norm = L._norm
_find = L.find_target_leg_safe


@dataclass(frozen=True)
class TradeTemplate:
    action: ActionId
    trade_kind: str
    params: tuple
    label: str


DEFAULT_HOLD = 5
DEFAULT_DTE = 35
WING_WIDTH = 5.0

# Action → structure parameters
ACTION_TEMPLATES: dict[ActionId, TradeTemplate | None] = {
    "cash": None,
    "iron_condor_neutral": TradeTemplate(
        "iron_condor_neutral",
        "ic",
        (DEFAULT_DTE, -0.16, 0.16, WING_WIDTH),
        "Delta-neutral iron condor",
    ),
    "iron_condor_neg_delta": TradeTemplate(
        "iron_condor_neg_delta",
        "ic_skew",
        (DEFAULT_DTE, -0.14, 0.22, WING_WIDTH),
        "Skewed iron condor (more call premium)",
    ),
    "iron_condor_pos_delta": TradeTemplate(
        "iron_condor_pos_delta",
        "ic_skew",
        (DEFAULT_DTE, -0.22, 0.14, WING_WIDTH),
        "Skewed iron condor (more put premium)",
    ),
    "put_credit_spread": TradeTemplate(
        "put_credit_spread",
        "vert",
        (DEFAULT_DTE, -0.20, WING_WIDTH),
        "Put credit spread",
    ),
    "call_debit_spread": TradeTemplate(
        "call_debit_spread",
        "call_debit",
        (DEFAULT_DTE, 0.35, WING_WIDTH),
        "Bull call debit spread",
    ),
    "put_debit_spread": TradeTemplate(
        "put_debit_spread",
        "put_debit",
        (DEFAULT_DTE, -0.35, WING_WIDTH),
        "Bear put debit spread",
    ),
    "long_call_calendar": TradeTemplate(
        "long_call_calendar",
        "call_cal",
        (14, 35, 0.30),
        "Long call calendar (near short / far long)",
    ),
    "long_straddle": TradeTemplate(
        "long_straddle",
        "sl",
        (DEFAULT_DTE,),
        "Long ATM straddle",
    ),
}


def iron_condor_pnl(
    entry_chain: OptionChain,
    exit_chain: OptionChain,
    exit_spy: float,
    target_dte: int,
    put_delta: float,
    call_delta: float,
    wing_width: float,
) -> float | None:
    """Short put spread + short call spread, same expiry."""
    for wing in (wing_width, wing_width * 2.0, wing_width * 3.0):
        put_pnl = L.short_vertical_put_pnl(
            entry_chain, exit_chain, exit_spy, target_dte, put_delta, wing
        )
        call_pnl = L.short_vertical_call_pnl(
            entry_chain, exit_chain, exit_spy, target_dte, call_delta, wing
        )
        if put_pnl is not None and call_pnl is not None:
            return float(put_pnl + call_pnl)
    return L.short_strangle_pnl(
        entry_chain, exit_chain, exit_spy, target_dte, put_delta, call_delta
    )


def can_build_trade(kind: str, chain: OptionChain, params: tuple) -> bool:
    """Entry-time probe: skip sessions where legs cannot be formed."""
    if not chain.contracts:
        return False
    if kind == "vert":
        dte, put_d, _wing = int(params[0]), float(params[1]), float(params[2])
        return _find(chain, dte, put_d, "P") is not None
    if kind == "sl":
        dte = int(params[0])
        call_k = [float(c.strike) for c in chain.contracts if str(c.option_type).upper() == "C"]
        put_k = [float(c.strike) for c in chain.contracts if str(c.option_type).upper() == "P"]
        if not call_k or not put_k:
            return False
        spot = float(sorted(set(call_k) & set(put_k))[len(set(call_k) & set(put_k)) // 2]) if (set(call_k) & set(put_k)) else (call_k[len(call_k) // 2] + put_k[len(put_k) // 2]) / 2.0
        return L.find_atm_straddle(chain, spot, dte) is not None
    if kind in ("ic", "ic_skew"):
        dte, put_d, call_d, wing = int(params[0]), float(params[1]), float(params[2]), float(params[3])
        ps = _find(chain, dte, put_d, "P")
        cs = _find(chain, dte, call_d, "C")
        if ps is None or cs is None:
            return False
        exp = _norm(ps.expiration)
        for w in (wing, wing * 2.0, wing * 3.0):
            if (
                find_contract_in_chain(chain, exp, float(ps.strike) - w, "P") is not None
                and find_contract_in_chain(chain, exp, float(cs.strike) + w, "C") is not None
            ):
                return True
        return True  # strangle fallback only needs ps/cs
    if kind == "call_debit":
        dte, long_d, wing = int(params[0]), float(params[1]), float(params[2])
        leg = _find(chain, dte, long_d, "C")
        return leg is not None and find_contract_in_chain(
            chain, _norm(leg.expiration), float(leg.strike) + wing, "C"
        ) is not None
    if kind == "put_debit":
        dte, long_d, wing = int(params[0]), float(params[1]), float(params[2])
        leg = _find(chain, dte, long_d, "P")
        return leg is not None and find_contract_in_chain(
            chain, _norm(leg.expiration), float(leg.strike) - wing, "P"
        ) is not None
    if kind == "call_cal":
        near_dte, far_dte, call_d = int(params[0]), int(params[1]), float(params[2])
        near = _find(chain, near_dte, call_d, "C")
        if near is None:
            return False
        return _pick_far_call_same_strike(chain, float(near.strike), far_dte, _norm(near.expiration)) is not None
    return False


def long_vertical_call_pnl(
    entry_chain: OptionChain,
    exit_chain: OptionChain,
    exit_spy: float,
    target_dte: int,
    long_delta: float,
    wing_width: float,
) -> float | None:
    """Bull call debit: buy lower-delta (OTM) call, sell higher strike call."""
    long_leg = _find(entry_chain, target_dte, long_delta, "C")
    if long_leg is None:
        return None
    exp = _norm(long_leg.expiration)
    short_k = float(long_leg.strike) + float(wing_width)
    short_leg = find_contract_in_chain(entry_chain, exp, short_k, "C")
    if short_leg is None:
        return None
    cost = (L._long_open_px(long_leg) - L._short_open_px(short_leg)) * MULT
    l_x = find_contract_in_chain(exit_chain, exp, float(long_leg.strike), "C")
    s_x = find_contract_in_chain(exit_chain, exp, short_k, "C")
    exit_d = _norm(exit_chain.as_of)
    if l_x is not None and s_x is not None:
        val = (L._long_close_px(l_x) - L._short_close_px(s_x)) * MULT
        return float(val - cost)
    if exit_d >= exp:
        val = (L._settle_long(long_leg, exit_spy) - L._settle_short(short_leg, exit_spy)) * MULT
        return float(val - cost)
    return None


def long_vertical_put_pnl(
    entry_chain: OptionChain,
    exit_chain: OptionChain,
    exit_spy: float,
    target_dte: int,
    long_delta: float,
    wing_width: float,
) -> float | None:
    """Bear put debit: buy higher-delta put, sell lower strike put."""
    long_leg = _find(entry_chain, target_dte, long_delta, "P")
    if long_leg is None:
        return None
    exp = _norm(long_leg.expiration)
    short_k = float(long_leg.strike) - float(wing_width)
    short_leg = find_contract_in_chain(entry_chain, exp, short_k, "P")
    if short_leg is None:
        return None
    cost = (L._long_open_px(long_leg) - L._short_open_px(short_leg)) * MULT
    l_x = find_contract_in_chain(exit_chain, exp, float(long_leg.strike), "P")
    s_x = find_contract_in_chain(exit_chain, exp, short_k, "P")
    exit_d = _norm(exit_chain.as_of)
    if l_x is not None and s_x is not None:
        val = (L._long_close_px(l_x) - L._short_close_px(s_x)) * MULT
        return float(val - cost)
    if exit_d >= exp:
        val = (L._settle_long(long_leg, exit_spy) - L._settle_short(short_leg, exit_spy)) * MULT
        return float(val - cost)
    return None


def long_call_calendar_pnl(
    entry_chain: OptionChain,
    exit_chain: OptionChain,
    exit_spy: float,
    near_dte: int,
    far_dte: int,
    call_delta: float,
) -> float | None:
    """Long calendar: sell near call, buy far call at same strike."""
    near = _find(entry_chain, near_dte, call_delta, "C")
    if near is None:
        return None
    strike = float(near.strike)
    far = _pick_far_call_same_strike(entry_chain, strike, far_dte, _norm(near.expiration))
    if far is None:
        return None
    net = (L._long_open_px(far) - L._short_open_px(near)) * MULT
    near_x = find_contract_in_chain(exit_chain, _norm(near.expiration), strike, "C")
    far_x = find_contract_in_chain(exit_chain, _norm(far.expiration), strike, "C")
    exit_d = _norm(exit_chain.as_of)
    if near_x is not None and far_x is not None:
        val = (L._long_close_px(far_x) - L._short_close_px(near_x)) * MULT
        return float(val - net)
    if exit_d >= _norm(near.expiration):
        near_settle = L._settle_short(near, exit_spy) * MULT
        if far_x is not None:
            far_val = L._long_close_px(far_x) * MULT
        elif exit_d >= _norm(far.expiration):
            far_val = L._settle_long(far, exit_spy) * MULT
        else:
            return None
        return float(far_val - near_settle - net)
    return None


def _pick_far_call_same_strike(
    chain: OptionChain,
    strike: float,
    target_far_dte: int,
    near_exp: pd.Timestamp,
) -> OptionContract | None:
    as_of = _norm(chain.as_of)
    best: tuple[int, OptionContract] | None = None
    for c in chain.contracts:
        if str(c.option_type).upper() != "C":
            continue
        if abs(float(c.strike) - strike) > 0.01:
            continue
        exp = _norm(c.expiration)
        if exp <= near_exp:
            continue
        dte = int((exp - as_of).days)
        err = abs(dte - int(target_far_dte))
        if best is None or err < best[0]:
            best = (err, c)
    return best[1] if best else None


TradeFn = Callable[[OptionChain, OptionChain, float, float, tuple], float | None]


def trade_fn_for_kind(kind: str) -> TradeFn:
    if kind == "ic":
        return lambda ch0, ch1, _se, sx, tp: iron_condor_pnl(
            ch0, ch1, sx, int(tp[0]), float(tp[1]), float(tp[2]), float(tp[3])
        )
    if kind == "ic_skew":
        return lambda ch0, ch1, _se, sx, tp: iron_condor_pnl(
            ch0, ch1, sx, int(tp[0]), float(tp[1]), float(tp[2]), float(tp[3])
        )
    if kind == "vert":
        return lambda ch0, ch1, _se, sx, tp: L.short_vertical_put_pnl(
            ch0, ch1, sx, int(tp[0]), float(tp[1]), float(tp[2])
        )
    if kind == "call_debit":
        return lambda ch0, ch1, _se, sx, tp: long_vertical_call_pnl(
            ch0, ch1, sx, int(tp[0]), float(tp[1]), float(tp[2])
        )
    if kind == "put_debit":
        return lambda ch0, ch1, _se, sx, tp: long_vertical_put_pnl(
            ch0, ch1, sx, int(tp[0]), float(tp[1]), float(tp[2])
        )
    if kind == "call_cal":
        return lambda ch0, ch1, _se, sx, tp: long_call_calendar_pnl(
            ch0, ch1, sx, int(tp[0]), int(tp[1]), float(tp[2])
        )
    if kind == "sl":
        return lambda ch0, ch1, se, sx, tp: L.long_straddle_pnl(ch0, ch1, se, sx, int(tp[0]))
    raise ValueError(f"Unknown trade kind: {kind}")
