"""
gex.py — implementazione dell'esposizione al gamma dei dealer.

Implementazione di riferimento della monografia «Gamma Exposure»
(Monografie Core Matrix, quaderno n. 1).

Tre principi guidano il disegno, e sono quelli discussi nel volume:

1. le greche si calcolano SENZA segno di posizione; il segno di inventario
   si applica a valle, dove puo' essere cambiato e messo alla prova;
2. il tempo residuo dipende da una convenzione dichiarata, non implicita;
3. lo zero gamma e' la RADICE del profilo GEX(x), con le greche ricalcolate
   a ogni prezzo ipotetico x. Non e' il punto in cui una somma cumulata a
   spot fisso cambia segno: sono due grandezze diverse.

Nessuna dipendenza esterna. Licenza MIT.

    from gex import Linea, Catena
    catena = Catena([...])
    catena.gex(spot=6508)          # dollari di delta per movimento dell'1%
    catena.zero_gamma()            # cambi di segno rilevati nell'intervallo
"""

from __future__ import annotations

import math
from dataclasses import dataclass, field
from typing import Iterable, Sequence

__all__ = ["Linea", "Catena", "greche", "beta", "moltiplicatore", "size_massima",
           "movimento_pareggio", "ORE_ANNO_BORSA", "ORE_ANNO_SOLARE"]
__version__ = "1.1.0"

SQRT2PI = math.sqrt(2.0 * math.pi)

ORE_ANNO_BORSA = 252 * 6.5      # 1638 — convenzione sulle ore di negoziazione
ORE_ANNO_SOLARE = 365 * 24      # 8760 — convenzione sulle ore di calendario


def _phi(x: float) -> float:
    return math.exp(-0.5 * x * x) / SQRT2PI


def _N(x: float) -> float:
    return 0.5 * (1.0 + math.erf(x / math.sqrt(2.0)))


def greche(S, K, T, sigma, r=0.0, q=0.0, tipo="call") -> dict:
    """Greche per UNA unita' di sottostante, senza segno di posizione.

    Charm e colour sono derivate rispetto al tempo CORRENTE t, non al tempo
    residuo T: un theta negativo significa che il valore cala col passare del
    tempo. Molte fonti usano la convenzione opposta.

    Il parametro q e' usato davvero in tutte le greche, charm e vanna
    compresi: le forme semplificate valide solo per q = 0 sono una fonte
    ricorrente di errori del 10% e oltre.
    """
    if not all(math.isfinite(v) for v in (S, K, T, sigma, r, q)):
        raise ValueError("i parametri numerici devono essere finiti")
    if S <= 0.0 or K <= 0.0:
        raise ValueError("spot e strike devono essere positivi")
    if T <= 0.0:
        raise ValueError("tempo residuo non positivo: filtrare le linee scadute")
    if sigma <= 0.0:
        raise ValueError("volatilita' non positiva")
    if tipo not in {"call", "put"}:
        raise ValueError("tipo deve essere 'call' oppure 'put'")

    v = sigma * math.sqrt(T)
    d1 = (math.log(S / K) + (r - q + 0.5 * sigma * sigma) * T) / v
    d2 = d1 - v
    dq, dr = math.exp(-q * T), math.exp(-r * T)

    delta = dq * (_N(d1) if tipo == "call" else _N(d1) - 1.0)
    gamma = dq * _phi(d1) / (S * v)
    vega = S * dq * _phi(d1) * math.sqrt(T)
    # charm: il termine q*delta distingue call e put, che per q > 0 differiscono
    charm = q * delta - dq * _phi(d1) * ((r - q) / v - d2 / (2.0 * T))
    vanna = -dq * _phi(d1) * d2 / sigma
    speed = -gamma / S * (1.0 + d1 / v)
    vomma = vega * d1 * d2 / sigma
    colour = gamma / (2.0 * T) * (
        1.0 + 2.0 * q * T + d1 * (2.0 * (r - q) * T - d2 * v) / v
    )

    if tipo == "call":
        theta = (-S * dq * _phi(d1) * sigma / (2.0 * math.sqrt(T))
                 - r * K * dr * _N(d2) + q * S * dq * _N(d1))
    else:
        theta = (-S * dq * _phi(d1) * sigma / (2.0 * math.sqrt(T))
                 + r * K * dr * _N(-d2) - q * S * dq * _N(-d1))

    return dict(d1=d1, d2=d2, delta=delta, gamma=gamma, vega=vega, theta=theta,
                charm=charm, vanna=vanna, speed=speed, vomma=vomma, colour=colour)


@dataclass
class Linea:
    """Una riga della catena.

    `oi_stock` e' l'open interest ereditato dalla chiusura precedente,
    `oi_flow` la stima dell'inventario aperto nella giornata corrente. Il
    volume insiste sulla distinzione perché i due termini hanno significati
    economici diversi. `oi_flow` resta una stima quando non sono disponibili
    dati intraday completi sulle nuove posizioni.
    """
    strike: float
    tipo: str                      # "call" | "put"
    ore_residue: float
    iv: float
    oi_stock: int = 0
    oi_flow: int = 0
    moltiplicatore: int = 100
    scadenza: str = ""

    @property
    def oi(self) -> int:
        return self.oi_stock + self.oi_flow


@dataclass
class Catena:
    """Una catena di opzioni, con la propria convenzione temporale."""
    linee: Sequence[Linea]
    convenzione: str = "borsa"     # "borsa" | "solare"
    r: float = 0.0
    q: float = 0.0
    w_call: float = +1.0           # ipotesi di segno: H0 = (+1, -1)
    w_put: float = -1.0

    def T(self, linea: Linea) -> float:
        if self.convenzione not in {"borsa", "solare"}:
            raise ValueError("convenzione deve essere 'borsa' oppure 'solare'")
        ore = ORE_ANNO_BORSA if self.convenzione == "borsa" else ORE_ANNO_SOLARE
        return linea.ore_residue / ore

    def peso(self, linea: Linea) -> float:
        if linea.tipo not in {"call", "put"}:
            raise ValueError("tipo della linea deve essere 'call' oppure 'put'")
        return self.w_call if linea.tipo == "call" else self.w_put

    # ------------------------------------------------------------ aggregati
    def _agg(self, spot: float, greca: str, scala, usa_flow=True) -> float:
        tot = 0.0
        for l in self.linee:
            g = greche(spot, l.strike, self.T(l), l.iv, self.r, self.q, l.tipo)[greca]
            oi = l.oi if usa_flow else l.oi_stock
            tot += self.peso(l) * g * oi * l.moltiplicatore * scala(spot, l)
        return tot

    def gex(self, spot: float, usa_flow=True) -> float:
        """Dollari di delta generati da un movimento dell'1% del sottostante."""
        return self._agg(spot, "gamma", lambda S, l: S * S * 0.01, usa_flow)

    def dex(self, spot: float) -> float:
        """Dollari di delta complessivi: il LIVELLO, non la sensibilita'."""
        return self._agg(spot, "delta", lambda S, l: S)

    def vex(self, spot: float) -> float:
        """Dollari di delta per un punto di volatilita' implicita."""
        return self._agg(spot, "vanna", lambda S, l: S * 0.01)

    def cex(self, spot: float, ore: float = 1.0) -> float:
        """Dollari di delta generati dal solo trascorrere di `ore` ore.

        Segno negativo = acquisti dei dealer.
        """
        anno = ORE_ANNO_BORSA if self.convenzione == "borsa" else ORE_ANNO_SOLARE
        return self._agg(spot, "charm", lambda S, l: S * (ore / anno))

    def gamma_efficace(self, spot: float, c_spotvol: float = 1.2) -> float:
        """Esposizione gamma-vanna sotto una relazione spot-vol specificata.

        `c_spotvol` e' il coefficiente c nell'ipotesi d_sigma = -c dS/S;
        non e' un coefficiente di correlazione. Il segno e l'entita' della
        correzione dipendono da c e dall'esposizione aggregata alla vanna.
        """
        return self.gex(spot) - c_spotvol * self.vex(spot)

    # --------------------------------------------------------- decomposizioni
    def per_strike(self, spot: float) -> dict:
        out: dict = {}
        for l in self.linee:
            g = greche(spot, l.strike, self.T(l), l.iv, self.r, self.q, l.tipo)["gamma"]
            v = self.peso(l) * g * l.oi * l.moltiplicatore * spot * spot * 0.01
            out[l.strike] = out.get(l.strike, 0.0) + v
        return out

    def per_scadenza(self, spot: float) -> dict:
        out: dict = {}
        for l in self.linee:
            g = greche(spot, l.strike, self.T(l), l.iv, self.r, self.q, l.tipo)["gamma"]
            v = self.peso(l) * g * l.oi * l.moltiplicatore * spot * spot * 0.01
            out[l.scadenza] = out.get(l.scadenza, 0.0) + v
        return out

    def sopra_sotto(self, spot: float) -> tuple:
        """GEX degli strike <= spot e > spot, restituiti separatamente."""
        giu = su = 0.0
        for l in self.linee:
            g = greche(spot, l.strike, self.T(l), l.iv, self.r, self.q, l.tipo)["gamma"]
            v = self.peso(l) * g * l.oi * l.moltiplicatore * spot * spot * 0.01
            if l.strike <= spot:
                giu += v
            else:
                su += v
        return giu, su

    def stock_flow(self, spot: float) -> tuple:
        tot = self.gex(spot, usa_flow=True)
        stock = self.gex(spot, usa_flow=False)
        return stock, tot - stock

    def concentrazione(self, spot: float) -> float:
        """Indice di Herfindahl sui pesi di gamma per strike.

        Piu' robusto dell'identificazione del singolo «muro»: risponde alla
        domanda operativa (quanto e' localizzata l'esposizione?) senza dover
        scommettere su uno strike.
        """
        per = self.per_strike(spot)
        tot = sum(abs(v) for v in per.values())
        if tot == 0.0:
            return 0.0
        return sum((abs(v) / tot) ** 2 for v in per.values())

    # -------------------------------------------------------------- profilo
    def profilo(self, x: float) -> float:
        """GEX(x): le greche vanno RICALCOLATE nello spot ipotetico x."""
        return self.gex(x)

    def zero_gamma(self, x_min=None, x_max=None, passo=1.0) -> list:
        """Radici associate ai cambi di segno rilevati nell'intervallo.

        Il profilo e' campionato con spaziatura `passo` e ogni cambio di segno
        e' raffinato per bisezione. Radici tangenti o coppie di radici piu'
        vicine del passo possono non essere rilevate. Se gli estremi non sono
        specificati, l'intervallo e' centrato sulla media degli strike e si
        estende per tre volatilita' giornaliere, usando la IV media.
        """
        if not self.linee:
            raise ValueError("la catena non contiene linee")
        if not math.isfinite(passo) or passo <= 0.0:
            raise ValueError("passo deve essere positivo e finito")
        if x_min is None or x_max is None:
            spot = sum(l.strike for l in self.linee) / len(self.linee)
            iv = sum(l.iv for l in self.linee) / len(self.linee)
            m = 3.0 * iv * math.sqrt(1.0 / 252.0)
            x_min = x_min if x_min is not None else max(spot * (1 - m), spot * 1e-6)
            x_max = x_max if x_max is not None else spot * (1 + m)
        if not (math.isfinite(x_min) and math.isfinite(x_max)):
            raise ValueError("gli estremi devono essere finiti")
        if not (0.0 < x_min < x_max):
            raise ValueError("richiesto 0 < x_min < x_max")
        radici, x = [], float(x_min)
        f = self.profilo(x)
        while x < x_max:
            y = min(x + passo, x_max)
            g = self.profilo(y)
            if f == 0.0:
                if not radici or abs(x - radici[-1]) > passo * 0.5:
                    radici.append(x)
            elif f * g < 0.0:
                lo, hi = x, y
                for _ in range(60):
                    mid = 0.5 * (lo + hi)
                    if self.profilo(lo) * self.profilo(mid) <= 0.0:
                        hi = mid
                    else:
                        lo = mid
                radici.append(0.5 * (lo + hi))
            x, f = y, g
        return radici


# --------------------------------------------------------------------------
# retroazione e dimensionamento
# --------------------------------------------------------------------------
def beta(gex_dollari, volume_dollari, sigma_giornaliera,
         orizzonte_min=None, minuti_seduta=390.0) -> float:
    """Parametro di retroazione del modello, riscalabile sull'orizzonte.

    beta = (sigma_d / 1%) * GEX / V$ ,  poi  beta(tau) = beta_d sqrt(tau_d/tau).
    Il riscalamento 1/sqrt(tau) e' un'ipotesi del modello; confronti tra
    orizzonti richiedono la stessa convenzione per volatilita' e durata.
    """
    if not math.isfinite(gex_dollari):
        raise ValueError("gex_dollari deve essere finito")
    if not (math.isfinite(volume_dollari) and volume_dollari > 0.0):
        raise ValueError("volume_dollari deve essere positivo e finito")
    if not (math.isfinite(sigma_giornaliera) and sigma_giornaliera >= 0.0):
        raise ValueError("sigma_giornaliera deve essere non negativa e finita")
    b = (sigma_giornaliera / 0.01) * gex_dollari / volume_dollari
    if orizzonte_min is None:
        return b
    if not (math.isfinite(orizzonte_min) and orizzonte_min > 0.0 and
            math.isfinite(minuti_seduta) and minuti_seduta > 0.0):
        raise ValueError("orizzonte e durata della seduta devono essere positivi e finiti")
    return b * math.sqrt(minuti_seduta / orizzonte_min)


def moltiplicatore(b: float) -> float:
    """1/(1+beta). Oltre beta = -1 il modello esce dal proprio dominio."""
    if not math.isfinite(b):
        raise ValueError("beta deve essere finito")
    if b <= -1.0:
        return float("inf")
    return 1.0 / (1.0 + b)


def size_massima(gamma_unitario, S, budget, k_sigma=3.0, x_shock=0.01,
                 sigma=0.12, T=None, moltiplicatore_contratto=100) -> int:
    """Contratti compatibili con il budget nei due scenari implementati.

    La perdita e' approssimata dal termine quadratico a gamma costante. Sono
    valutati uno spostamento di `k_sigma` deviazioni standard residue e uno
    shock relativo fisso; viene usato lo scenario piu' severo. Non sono inclusi
    gli altri fattori di rischio della posizione.
    """
    valori = (gamma_unitario, S, budget, k_sigma, x_shock, sigma,
              moltiplicatore_contratto)
    if not all(math.isfinite(v) for v in valori):
        raise ValueError("i parametri devono essere finiti")
    if not (S > 0.0 and budget >= 0.0 and k_sigma >= 0.0 and x_shock >= 0.0 and
            sigma >= 0.0 and moltiplicatore_contratto > 0.0):
        raise ValueError("parametri di dimensionamento fuori dominio")
    if T is not None and (not math.isfinite(T) or T < 0.0):
        raise ValueError("T deve essere non negativo e finito")
    g_dollari = abs(gamma_unitario) * moltiplicatore_contratto * S * S
    x_sd = k_sigma * sigma * math.sqrt(T) if T is not None else 0.0
    perdita = 0.5 * g_dollari * max(x_sd ** 2, x_shock ** 2)
    return int(math.floor(budget / perdita)) if perdita > 0 else 0


def movimento_pareggio(S, sigma, T) -> float:
    """Pareggio gamma-theta locale con r=q=0: S sigma sqrt(T).

    La relazione usa l'approssimazione quadratica del P&L coperto in delta.
    """
    if not all(math.isfinite(v) for v in (S, sigma, T)):
        raise ValueError("i parametri devono essere finiti")
    if not (S > 0.0 and sigma >= 0.0 and T >= 0.0):
        raise ValueError("richiesti S > 0, sigma >= 0 e T >= 0")
    return S * sigma * math.sqrt(T)
