"""
Black-Scholes-Merton Option Pricing Model.

Implements European option pricing with crypto adaptations:
- Full Greeks calculation (Delta, Gamma, Vega, Theta, Rho)
- Implied volatility via Newton-Raphson
- Staking yield adjustment (for ETH)

Mathematical Foundation:
- BSM PDE: ∂V/∂t + rS∂V/∂S + ½σ²S²∂²V/∂S² - rV = 0
- Call: C = SN(d₊) - Ke^{-rτ}N(d₋)
- Put:  P = Ke^{-rτ}N(-d₋) - SN(-d₊)

where:
    d± = [ln(S/K) + (r ± σ²/2)τ] / (σ√τ)

Greeks:
- Delta: ∂V/∂S = N(d₊) for call
- Gamma: ∂²V/∂S² = n(d₊)/(Sσ√τ)
- Vega:  ∂V/∂σ = S√τ n(d₊)
- Theta: ∂V/∂t = -Sσn(d₊)/(2√τ) - rKe^{-rτ}N(d₋)
- Rho:   ∂V/∂r = Kτe^{-rτ}N(d₋)

Crypto considerations:
- High implied vol: 20%-150% vs equity 15%-25%
- Staking yield q: F = Se^{(r-q)T} for ETH (~3-5% annually)
- 24/7 trading: τ = T/365 (no business day adjustment)
"""
import math
from dataclasses import dataclass
from enum import Enum


class OptionType(Enum):
    CALL = "call"
    PUT = "put"


@dataclass(slots=True)
class BSMGreeks:
    """Container for option Greeks."""
    delta: float      # ∂V/∂S - price sensitivity
    gamma: float      # ∂²V/∂S² - delta sensitivity
    vega: float       # ∂V/∂σ - volatility sensitivity (per 1% move)
    theta: float      # ∂V/∂t - time decay (per day)
    rho: float        # ∂V/∂r - rate sensitivity (per 1% move)
    # Second-order Greeks
    vanna: float      # ∂²V/∂S∂σ - delta-vol cross
    volga: float      # ∂²V/∂σ² - vega convexity


class BSMPricer:
    """
    Black-Scholes-Merton option pricer with full Greeks.

    Handles European calls and puts with optional dividend yield
    (for staking assets like ETH).

    All methods are static for stateless computation.
    """

    @staticmethod
    def _norm_cdf(x: float) -> float:
        """
        Standard normal CDF approximation.

        Uses Abramowitz & Stegun approximation (7.1.26).
        Accurate to 7.5e-8.
        """
        # Constants
        a1 = 0.254829592
        a2 = -0.284496736
        a3 = 1.421413741
        a4 = -1.453152027
        a5 = 1.061405429
        p = 0.3275911

        sign = 1 if x >= 0 else -1
        x = abs(x) / math.sqrt(2)

        t = 1.0 / (1.0 + p * x)
        y = 1.0 - (((((a5 * t + a4) * t) + a3) * t + a2) * t + a1) * t * math.exp(-x * x)

        return 0.5 * (1.0 + sign * y)

    @staticmethod
    def _norm_pdf(x: float) -> float:
        """Standard normal PDF."""
        return math.exp(-0.5 * x * x) / math.sqrt(2 * math.pi)

    @staticmethod
    def _d_plus_minus(
        S: float,
        K: float,
        T: float,
        r: float,
        sigma: float,
        q: float = 0.0,
    ) -> tuple[float, float]:
        """
        Calculate d+ and d- for BSM formula.

        d± = [ln(S/K) + (r - q ± σ²/2)T] / (σ√T)
        """
        if T <= 0 or sigma <= 0:
            return 0.0, 0.0

        sqrt_T = math.sqrt(T)
        d_plus = (math.log(S / K) + (r - q + 0.5 * sigma * sigma) * T) / (sigma * sqrt_T)
        d_minus = d_plus - sigma * sqrt_T

        return d_plus, d_minus

    @staticmethod
    def price_call(
        S: float,
        K: float,
        T: float,
        r: float,
        sigma: float,
        q: float = 0.0,
    ) -> float:
        """
        Price a European call option.

        C = Se^{-qT}N(d₊) - Ke^{-rT}N(d₋)

        Args:
            S: Spot price
            K: Strike price
            T: Time to expiry (years)
            r: Risk-free rate (annualized)
            sigma: Volatility (annualized)
            q: Dividend/staking yield (for ETH, ~0.03-0.05)

        Returns:
            Call option price
        """
        if T <= 0:
            return max(S - K, 0.0)

        d_plus, d_minus = BSMPricer._d_plus_minus(S, K, T, r, sigma, q)

        call = (S * math.exp(-q * T) * BSMPricer._norm_cdf(d_plus) -
                K * math.exp(-r * T) * BSMPricer._norm_cdf(d_minus))

        return max(call, 0.0)

    @staticmethod
    def price_put(
        S: float,
        K: float,
        T: float,
        r: float,
        sigma: float,
        q: float = 0.0,
    ) -> float:
        """
        Price a European put option.

        P = Ke^{-rT}N(-d₋) - Se^{-qT}N(-d₊)

        Can also use put-call parity: P = C - Se^{-qT} + Ke^{-rT}
        """
        if T <= 0:
            return max(K - S, 0.0)

        d_plus, d_minus = BSMPricer._d_plus_minus(S, K, T, r, sigma, q)

        put = (K * math.exp(-r * T) * BSMPricer._norm_cdf(-d_minus) -
               S * math.exp(-q * T) * BSMPricer._norm_cdf(-d_plus))

        return max(put, 0.0)

    @staticmethod
    def price(
        option_type: OptionType,
        S: float,
        K: float,
        T: float,
        r: float,
        sigma: float,
        q: float = 0.0,
    ) -> float:
        """Price call or put based on option type."""
        if option_type == OptionType.CALL:
            return BSMPricer.price_call(S, K, T, r, sigma, q)
        else:
            return BSMPricer.price_put(S, K, T, r, sigma, q)

    @staticmethod
    def greeks(
        option_type: OptionType,
        S: float,
        K: float,
        T: float,
        r: float,
        sigma: float,
        q: float = 0.0,
    ) -> BSMGreeks:
        """
        Calculate all Greeks for an option.

        Returns Greeks with conventional scaling:
        - Vega: per 1% vol move (multiply by 0.01)
        - Theta: per day (multiply by 1/365)
        - Rho: per 1% rate move (multiply by 0.01)
        """
        if T <= 0 or sigma <= 0:
            # At expiry, only delta matters
            intrinsic = S - K if option_type == OptionType.CALL else K - S
            return BSMGreeks(
                delta=1.0 if intrinsic > 0 else 0.0,
                gamma=0.0, vega=0.0, theta=0.0, rho=0.0,
                vanna=0.0, volga=0.0,
            )

        d_plus, d_minus = BSMPricer._d_plus_minus(S, K, T, r, sigma, q)
        sqrt_T = math.sqrt(T)

        N_d_plus = BSMPricer._norm_cdf(d_plus)
        N_d_minus = BSMPricer._norm_cdf(d_minus)
        n_d_plus = BSMPricer._norm_pdf(d_plus)

        exp_qT = math.exp(-q * T)
        exp_rT = math.exp(-r * T)

        # Delta
        if option_type == OptionType.CALL:
            delta = exp_qT * N_d_plus
        else:
            delta = exp_qT * (N_d_plus - 1)

        # Gamma (same for call and put)
        gamma = exp_qT * n_d_plus / (S * sigma * sqrt_T)

        # Vega (same for call and put)
        # Standard vega (per unit vol), scale to per 1% later
        vega = S * exp_qT * sqrt_T * n_d_plus

        # Theta
        term1 = -S * sigma * exp_qT * n_d_plus / (2 * sqrt_T)
        if option_type == OptionType.CALL:
            term2 = -r * K * exp_rT * N_d_minus
            term3 = q * S * exp_qT * N_d_plus
        else:
            term2 = r * K * exp_rT * BSMPricer._norm_cdf(-d_minus)
            term3 = -q * S * exp_qT * BSMPricer._norm_cdf(-d_plus)
        theta = term1 + term2 + term3

        # Rho
        if option_type == OptionType.CALL:
            rho = K * T * exp_rT * N_d_minus
        else:
            rho = -K * T * exp_rT * BSMPricer._norm_cdf(-d_minus)

        # Second-order Greeks
        # Vanna = ∂Delta/∂sigma = -n(d+) * d- / sigma
        vanna = -exp_qT * n_d_plus * d_minus / sigma

        # Volga = ∂Vega/∂sigma = Vega * d+ * d- / sigma
        volga = vega * d_plus * d_minus / sigma

        return BSMGreeks(
            delta=delta,
            gamma=gamma,
            vega=vega * 0.01,        # Per 1% vol move
            theta=theta / 365,        # Per day
            rho=rho * 0.01,          # Per 1% rate move
            vanna=vanna * 0.01,      # Per 1% vol move
            volga=volga * 0.0001,    # Per 1% vol move squared
        )

    @staticmethod
    def implied_vol(
        market_price: float,
        option_type: OptionType,
        S: float,
        K: float,
        T: float,
        r: float,
        q: float = 0.0,
        tol: float = 1e-6,
        max_iter: int = 100,
    ) -> float | None:
        """
        Calculate implied volatility using Newton-Raphson.

        Newton-Raphson: σₙ₊₁ = σₙ - (C_BS(σₙ) - C_mkt) / Vega(σₙ)

        Quadratic convergence - typically 3-5 iterations.

        Args:
            market_price: Observed option price
            option_type: Call or put
            S, K, T, r, q: BSM parameters
            tol: Convergence tolerance
            max_iter: Maximum iterations

        Returns:
            Implied volatility or None if not converged
        """
        if T <= 0 or market_price <= 0:
            return None

        # Check intrinsic value bounds
        if option_type == OptionType.CALL:
            intrinsic = max(S * math.exp(-q * T) - K * math.exp(-r * T), 0)
            max_val = S * math.exp(-q * T)
        else:
            intrinsic = max(K * math.exp(-r * T) - S * math.exp(-q * T), 0)
            max_val = K * math.exp(-r * T)

        if market_price < intrinsic - 1e-10 or market_price > max_val + 1e-10:
            return None

        # Initial guess (use approximation)
        sigma = 0.5  # Start at 50% vol

        for _ in range(max_iter):
            # Price at current sigma
            price = BSMPricer.price(option_type, S, K, T, r, sigma, q)

            # Error
            diff = price - market_price

            if abs(diff) < tol:
                return sigma

            # Vega for Newton step
            d_plus, _ = BSMPricer._d_plus_minus(S, K, T, r, sigma, q)
            vega = S * math.exp(-q * T) * math.sqrt(T) * BSMPricer._norm_pdf(d_plus)

            if vega < 1e-10:
                # Vega too small, try bisection
                sigma *= 1.5 if diff < 0 else 0.7
                continue

            # Newton step
            sigma_new = sigma - diff / vega

            # Bound sigma to reasonable range
            sigma = max(0.01, min(5.0, sigma_new))

        # Did not converge - return best guess
        return sigma

    @staticmethod
    def forward_price(
        S: float,
        T: float,
        r: float,
        q: float = 0.0,
    ) -> float:
        """
        Calculate forward price.

        F = S × e^{(r-q)T}

        For crypto with staking: q = staking_yield
        """
        return S * math.exp((r - q) * T)


class BSMMonteCarloGreeks:
    """
    Monte Carlo Greeks calculation using finite differences.

    Useful for path-dependent options or model validation.
    Uses antithetic variates for variance reduction.
    """

    @staticmethod
    def delta_mc(
        option_type: OptionType,
        S: float,
        K: float,
        T: float,
        r: float,
        sigma: float,
        q: float = 0.0,
        bump: float = 0.01,
        n_paths: int = 10000,
    ) -> float:
        """
        Calculate delta using Monte Carlo with central difference.

        Δ ≈ (V(S+ε) - V(S-ε)) / (2ε)
        """
        import random

        def simulate_payoff(spot: float) -> float:
            """Simulate option payoff at expiry."""
            total = 0.0
            drift = (r - q - 0.5 * sigma * sigma) * T
            vol_term = sigma * math.sqrt(T)

            for _ in range(n_paths // 2):
                z = random.gauss(0, 1)

                # Antithetic pair
                S_T_1 = spot * math.exp(drift + vol_term * z)
                S_T_2 = spot * math.exp(drift - vol_term * z)

                if option_type == OptionType.CALL:
                    total += max(S_T_1 - K, 0) + max(S_T_2 - K, 0)
                else:
                    total += max(K - S_T_1, 0) + max(K - S_T_2, 0)

            return total / n_paths * math.exp(-r * T)

        eps = S * bump
        V_up = simulate_payoff(S + eps)
        V_down = simulate_payoff(S - eps)

        return (V_up - V_down) / (2 * eps)


# Convenience functions
def price_call(S, K, T, r, sigma, q=0.0):
    """Convenience function for call pricing."""
    return BSMPricer.price_call(S, K, T, r, sigma, q)


def price_put(S, K, T, r, sigma, q=0.0):
    """Convenience function for put pricing."""
    return BSMPricer.price_put(S, K, T, r, sigma, q)


def implied_vol(market_price, option_type, S, K, T, r, q=0.0):
    """Convenience function for implied vol calculation."""
    return BSMPricer.implied_vol(market_price, option_type, S, K, T, r, q)
