"""Data models for option chain results.

Uses :func:`dataclasses.dataclass` with ``slots=True`` for maximum memory
efficiency — critical when holding thousands of strike rows in memory.
"""

from __future__ import annotations

from dataclasses import dataclass, field
from datetime import datetime
from typing import Any


@dataclass(slots=True, frozen=True)
class OptionChainRow:
    """A single strike row in an option chain.

    Both ``call_*`` and ``put_*`` fields can be ``None`` when the exchange
    does not provide data for that side (e.g. deep ITM / OTM strikes with no
    open interest).

    Attributes:
        strike:    Strike price.
        call_oi:   Call open interest (contracts).
        call_coi:  Call change in open interest vs. previous session.
        put_oi:    Put open interest (contracts).
        put_coi:   Put change in open interest vs. previous session.
        call_iv:   Call implied volatility (%).
        put_iv:    Put implied volatility (%).
        call_ltp:  Call last traded price.
        put_ltp:   Put last traded price.
        call_vol:  Call traded volume (contracts), if available.
        put_vol:   Put traded volume (contracts), if available.
        call_bid:  Call best bid price, if available.
        put_bid:   Put best bid price, if available.
        call_ask:  Call best ask price, if available.
        put_ask:   Put best ask price, if available.
    """

    strike: float
    call_oi: int = 0
    call_coi: int = 0
    put_oi: int = 0
    put_coi: int = 0
    call_iv: float | None = None
    put_iv: float | None = None
    call_ltp: float | None = None
    put_ltp: float | None = None
    call_vol: int | None = None
    put_vol: int | None = None
    call_bid: float | None = None
    put_bid: float | None = None
    call_ask: float | None = None
    put_ask: float | None = None

    @property
    def pcr(self) -> float | None:
        """Put/Call ratio by open interest.

        Returns ``None`` when call OI is zero to avoid division by zero.
        """
        if self.call_oi == 0:
            return None
        return self.put_oi / self.call_oi

    @property
    def total_oi(self) -> int:
        """Sum of call and put open interest at this strike."""
        return self.call_oi + self.put_oi

    def to_dict(self) -> dict[str, Any]:
        """Serialize to a plain dict (JSON-serializable)."""
        return {
            "strike": self.strike,
            "call_oi": self.call_oi,
            "call_coi": self.call_coi,
            "put_oi": self.put_oi,
            "put_coi": self.put_coi,
            "call_iv": self.call_iv,
            "put_iv": self.put_iv,
            "call_ltp": self.call_ltp,
            "put_ltp": self.put_ltp,
            "call_vol": self.call_vol,
            "put_vol": self.put_vol,
            "call_bid": self.call_bid,
            "put_bid": self.put_bid,
            "call_ask": self.call_ask,
            "put_ask": self.put_ask,
            "pcr": self.pcr,
        }


@dataclass(slots=True)
class OptionChainResult:
    """Complete option chain result for a symbol.

    Attributes:
        symbol:      Trading symbol (e.g. ``"NIFTY"``).
        exchange:    Exchange name (``"NSE"`` or ``"BSE"``).
        expiry:      Nearest expiry date string as returned by the exchange.
        spot_price:  Underlying spot price at fetch time.
        atm_strike:  At-the-money strike (closest to spot_price).
        data:        List of :class:`OptionChainRow`, sorted by strike ascending.
        fetched_at:  UTC timestamp when this result was fetched.
        expiry_dates: All available expiry dates for the symbol (raw exchange strings).
    """

    symbol: str
    exchange: str
    expiry: str | None
    spot_price: float | None
    atm_strike: float | None
    data: list[OptionChainRow]
    fetched_at: datetime
    expiry_dates: list[str] = field(default_factory=list)

    # ── Derived properties ────────────────────────────────────────────────────

    @property
    def strikes(self) -> list[float]:
        """Sorted list of all available strike prices."""
        return [row.strike for row in self.data]

    @property
    def total_call_oi(self) -> int:
        """Sum of all call open interest across all strikes."""
        return sum(row.call_oi for row in self.data)

    @property
    def total_put_oi(self) -> int:
        """Sum of all put open interest across all strikes."""
        return sum(row.put_oi for row in self.data)

    @property
    def pcr(self) -> float | None:
        """Overall Put/Call ratio.

        Returns ``None`` if total call OI is zero.
        """
        total_call = self.total_call_oi
        if total_call == 0:
            return None
        return self.total_put_oi / total_call

    @property
    def max_pain_strike(self) -> float | None:
        """Max pain strike — the strike with the highest total OI.

        This is a simplified max pain calculation based on total open interest
        (call OI + put OI). A full max pain calculation requires pricing each
        option at expiry, which is beyond the scope of this model.
        """
        if not self.data:
            return None
        return max(self.data, key=lambda r: r.total_oi).strike

    def atm_window(self, n: int = 5) -> list[OptionChainRow]:
        """Return the ``n`` strikes above and below the ATM strike.

        Args:
            n: Number of strikes on each side of ATM.

        Returns:
            Up to ``2*n + 1`` rows centred on the ATM strike.
        """
        if self.atm_strike is None or not self.data:
            return list(self.data)
        strikes = self.strikes
        try:
            atm_idx = strikes.index(self.atm_strike)
        except ValueError:
            # ATM not exactly in list — find nearest
            atm_idx = min(range(len(strikes)), key=lambda i: abs(strikes[i] - self.atm_strike))  # type: ignore[operator]
        lo = max(0, atm_idx - n)
        hi = min(len(self.data), atm_idx + n + 1)
        return self.data[lo:hi]

    def to_dict(self) -> dict[str, Any]:
        """Serialize to a plain JSON-serializable dict."""
        return {
            "symbol": self.symbol,
            "exchange": self.exchange,
            "expiry": self.expiry,
            "spot_price": self.spot_price,
            "atm_strike": self.atm_strike,
            "pcr": self.pcr,
            "max_pain_strike": self.max_pain_strike,
            "total_call_oi": self.total_call_oi,
            "total_put_oi": self.total_put_oi,
            "fetched_at": self.fetched_at.isoformat(),
            "expiry_dates": self.expiry_dates,
            "data": [row.to_dict() for row in self.data],
        }


__all__ = ["OptionChainResult", "OptionChainRow"]
