"""NSE option chain response parser.

Transforms the raw JSON payload from NSE's ``/api/option-chain-v3``
endpoint into a typed :class:`~indiaopt.models.option_chain.OptionChainResult`.

NSE response schema (simplified)::

    {
        "records": {
            "expiryDates": ["27-Jun-2024", ...],
            "data": [
                {
                    "strikePrice": 22000,
                    "CE": {
                        "openInterest": 12345,
                        "changeinOpenInterest": 678,
                        "impliedVolatility": 12.34,
                        "lastPrice": 45.6,
                        ...
                    },
                    "PE": { ... }
                },
                ...
            ],
            "underlyingValue": 22150.35
        },
        "filtered": { ... }
    }
"""

from __future__ import annotations

import logging
from datetime import datetime, timezone
from typing import Any

from indiaopt.exceptions import ParseError
from indiaopt.models.option_chain import OptionChainResult, OptionChainRow
from indiaopt.utils.coerce import safe_float, safe_int, safe_int_or_zero

logger = logging.getLogger("indiaopt.parsers.nse")

EXCHANGE = "NSE"


def parse_option_chain(raw: dict[str, Any], symbol: str) -> OptionChainResult:
    """Parse an NSE ``option-chain-v3`` response into :class:`OptionChainResult`.

    Args:
        raw:    Raw JSON dict from NSE API.
        symbol: The trading symbol (used for error messages).

    Returns:
        Fully typed :class:`OptionChainResult`.

    Raises:
        :class:`~indiaopt.exceptions.ParseError`: If the response is missing
            required structure (e.g. ``"records"`` key absent).
    """
    try:
        return _parse(raw, symbol)
    except ParseError:
        raise
    except Exception as exc:
        raise ParseError(
            f"Unexpected error parsing NSE response for {symbol}: {exc}",
            symbol=symbol,
            exchange=EXCHANGE,
            raw_preview=str(raw)[:200],
        ) from exc


def _parse(raw: dict[str, Any], symbol: str) -> OptionChainResult:
    records = raw.get("records")
    if not isinstance(records, dict):
        raise ParseError(
            f"NSE response missing 'records' key for {symbol}. "
            f"Got type={type(raw.get('records')).__name__!r}. "
            "This often means the session cookie expired — the client will refresh.",
            symbol=symbol,
            exchange=EXCHANGE,
            raw_preview=str(raw)[:200],
        )

    # ── Expiry dates ──────────────────────────────────────────────────────────
    expiry_dates: list[str] = records.get("expiryDates") or []
    nearest_expiry: str | None = expiry_dates[0] if expiry_dates else None

    # ── Spot price ────────────────────────────────────────────────────────────
    spot_price: float | None = safe_float(records.get("underlyingValue"))

    # ── Strike rows ───────────────────────────────────────────────────────────
    raw_data: list[Any] = records.get("data") or []
    if not raw_data:
        logger.warning(
            "NSE response for %s has empty 'records.data'. "
            "Market may be closed or symbol not found.",
            symbol,
        )

    rows: list[OptionChainRow] = []
    skipped = 0
    for item in raw_data:
        if not isinstance(item, dict):
            skipped += 1
            continue
        strike = safe_float(item.get("strikePrice"))
        if strike is None:
            skipped += 1
            continue
        ce: dict[str, Any] = item.get("CE") or {}
        pe: dict[str, Any] = item.get("PE") or {}
        rows.append(
            OptionChainRow(
                strike=strike,
                call_oi=safe_int_or_zero(ce.get("openInterest")),
                call_coi=safe_int_or_zero(ce.get("changeinOpenInterest")),
                put_oi=safe_int_or_zero(pe.get("openInterest")),
                put_coi=safe_int_or_zero(pe.get("changeinOpenInterest")),
                call_iv=safe_float(ce.get("impliedVolatility")),
                put_iv=safe_float(pe.get("impliedVolatility")),
                call_ltp=safe_float(ce.get("lastPrice")),
                put_ltp=safe_float(pe.get("lastPrice")),
                call_vol=safe_int(ce.get("totalTradedVolume")),
                put_vol=safe_int(pe.get("totalTradedVolume")),
                call_bid=safe_float(ce.get("bidprice")),
                put_bid=safe_float(pe.get("bidprice")),
                call_ask=safe_float(ce.get("askPrice")),
                put_ask=safe_float(pe.get("askPrice")),
            )
        )

    if skipped:
        logger.debug("Skipped %d malformed rows for %s.", skipped, symbol)

    # Sort by strike ascending
    rows.sort(key=lambda r: r.strike)

    # ── ATM strike ────────────────────────────────────────────────────────────
    atm_strike: float | None = None
    if spot_price is not None and rows:
        atm_strike = min(rows, key=lambda r: abs(r.strike - spot_price)).strike
        logger.debug(
            "ATM strike for %s: %.2f (spot=%.2f, %d strikes available)",
            symbol,
            atm_strike,
            spot_price,
            len(rows),
        )

    return OptionChainResult(
        symbol=symbol,
        exchange=EXCHANGE,
        expiry=nearest_expiry,
        spot_price=spot_price,
        atm_strike=atm_strike,
        data=rows,
        fetched_at=datetime.now(timezone.utc),
        expiry_dates=expiry_dates,
    )


__all__ = ["parse_option_chain"]
