"""
Heston (1993) stochastic-volatility closed-form option pricing, from scratch.

    dS_t = r S_t dt + sqrt(v_t) S_t dW1
    dv_t = kappa(theta - v_t) dt + sigma sqrt(v_t) dW2,   corr(dW1,dW2) = rho
Parameters: v0 (initial var), kappa (mean-reversion speed), theta (long-run var),
            sigma (vol-of-vol), rho (leverage correlation, <0 for equity).

European call by the two-probability Fourier formula  C = S P1 - K e^{-rT} P2,
using the numerically stable "little Heston trap" characteristic function (Albrecher et al. 2007).
"""
import numpy as np
from scipy.stats import norm
from scipy.optimize import brentq

_trapz = getattr(np, "trapezoid", getattr(np, "trapz", None))


def _cf(phi, j, S, v0, r, T, kappa, theta, sigma, rho):
    """Heston characteristic function of ln S_T for j = 1, 2 (little-trap form)."""
    u = 0.5 if j == 1 else -0.5
    bj = (kappa - rho * sigma) if j == 1 else kappa
    a = kappa * theta
    rspi = rho * sigma * 1j * phi
    d = np.sqrt((rspi - bj) ** 2 - sigma ** 2 * (2 * u * 1j * phi - phi ** 2))
    g2 = (bj - rspi - d) / (bj - rspi + d)                       # = 1/g  (stable branch)
    ed = np.exp(-d * T)
    D = ((bj - rspi - d) / sigma ** 2) * (1 - ed) / (1 - g2 * ed)
    C = r * 1j * phi * T + (a / sigma ** 2) * ((bj - rspi - d) * T - 2 * np.log((1 - g2 * ed) / (1 - g2)))
    return np.exp(C + D * v0 + 1j * phi * np.log(S))


def heston_call(S, K, r, T, v0, kappa, theta, sigma, rho, phimax=200.0, nphi=2000):
    phi = np.linspace(1e-8, phimax, nphi)
    f1 = _cf(phi, 1, S, v0, r, T, kappa, theta, sigma, rho)
    f2 = _cf(phi, 2, S, v0, r, T, kappa, theta, sigma, rho)
    K = np.atleast_1d(np.asarray(K, float)); out = np.empty(len(K))
    for i, k in enumerate(K):
        e = np.exp(-1j * phi * np.log(k))
        P1 = 0.5 + (1 / np.pi) * _trapz((e * f1 / (1j * phi)).real, phi)
        P2 = 0.5 + (1 / np.pi) * _trapz((e * f2 / (1j * phi)).real, phi)
        out[i] = S * P1 - k * np.exp(-r * T) * P2
    return out if out.size > 1 else float(out[0])


def heston_mc(S, K, r, T, v0, kappa, theta, sigma, rho, M=200000, steps=250, seed=0):
    """Euler (full-truncation) Monte-Carlo cross-check; returns (price, std-err)."""
    dt = T / steps; rng = np.random.default_rng(seed)
    s = np.full(M, float(S)); v = np.full(M, float(v0))
    for _ in range(steps):
        z1 = rng.standard_normal(M); z2 = rho * z1 + np.sqrt(1 - rho ** 2) * rng.standard_normal(M)
        vp = np.maximum(v, 0.0)
        s *= np.exp((r - 0.5 * vp) * dt + np.sqrt(vp * dt) * z1)
        v = v + kappa * (theta - vp) * dt + sigma * np.sqrt(vp * dt) * z2                # full truncation
    disc = np.exp(-r * T); pay = np.maximum(s - float(K), 0.0)
    return disc * pay.mean(), disc * pay.std() / np.sqrt(M)


def bs_call(S, K, r, T, sig):
    if sig <= 0 or T <= 0: return max(S - K * np.exp(-r * T), 0.0)
    d1 = (np.log(S / K) + (r + 0.5 * sig ** 2) * T) / (sig * np.sqrt(T)); d2 = d1 - sig * np.sqrt(T)
    return S * norm.cdf(d1) - K * np.exp(-r * T) * norm.cdf(d2)


def bs_iv(price, S, K, r, T):
    if price <= max(S - K * np.exp(-r * T), 0.0) + 1e-10: return np.nan
    try:
        return brentq(lambda x: bs_call(S, K, r, T, x) - price, 1e-6, 5.0, maxiter=200)
    except Exception:
        return np.nan
