"""NSEClient — high-level async client for NSE option chain data.

Usage::

    import asyncio
    from indiaopt import NSEClient

    async def main():
        async with NSEClient() as client:
            result = await client.fetch_option_chain("NIFTY", is_index=True)
            print(f"Spot: {result.spot_price}, ATM: {result.atm_strike}")
            for row in result.atm_window(n=5):
                print(f"{row.strike:>8}  CE OI: {row.call_oi:>10}  PE OI: {row.put_oi:>10}")

    asyncio.run(main())
"""

from __future__ import annotations

import asyncio
import logging
from concurrent.futures import ThreadPoolExecutor
from typing import Any

from indiaopt.circuit import CircuitBreaker, get_default_breaker
from indiaopt.config.settings import Settings, get_settings
from indiaopt.constants import (
    NSE_BASE_URL,
    NSE_CONTRACT_INFO_PATH,
    NSE_OPTION_CHAIN_PATH,
    RATE_LIMIT_STATUS_CODES,
)
from indiaopt.exceptions import (
    FetchError,
    FetchTimeoutError,
    NetworkError,
    ParseError,
    RateLimitError,
)
from indiaopt.http.retry import RetryPolicy
from indiaopt.http.session import create_nse_session
from indiaopt.models.option_chain import OptionChainResult
from indiaopt.parsers import nse as nse_parser

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

EXCHANGE = "NSE"


class NSEClient:
    """Async client for NSE option chain and market data.

    Manages its own session lifetime, thread pool, circuit breaker, and
    retry policy. Supports both context-manager and direct usage.

    Args:
        settings: Optional :class:`~indiaopt.config.settings.Settings` override.
                  Uses global settings (from env/``.env``) when not provided.
        breaker:  Optional :class:`~indiaopt.circuit.CircuitBreaker` override.
                  Shares the global default when not provided.

    Context manager (recommended)::

        async with NSEClient() as client:
            result = await client.fetch_option_chain("NIFTY")

    Direct usage (must call :meth:`close` manually)::

        client = NSEClient()
        try:
            result = await client.fetch_option_chain("NIFTY")
        finally:
            await client.close()
    """

    def __init__(
        self,
        settings: Settings | None = None,
        breaker: CircuitBreaker | None = None,
    ) -> None:
        self._settings: Settings = settings or get_settings()
        self._breaker: CircuitBreaker = breaker or get_default_breaker(
            failure_threshold=self._settings.circuit_failure_threshold,
            backoff_seconds=self._settings.circuit_backoff_seconds,
        )
        self._retry_policy = RetryPolicy.from_settings(self._settings)
        self._pool = ThreadPoolExecutor(
            max_workers=self._settings.thread_pool_workers,
            thread_name_prefix="indiaopt-nse",
        )
        self._session: Any = None  # curl_cffi Session — created on demand

    # ── Context manager ───────────────────────────────────────────────────────

    async def __aenter__(self) -> NSEClient:
        return self

    async def __aexit__(self, *_: object) -> None:
        await self.close()

    async def close(self) -> None:
        """Release all resources (session, thread pool)."""
        if self._session is not None:
            try:
                self._session.close()
            except Exception:
                pass
            self._session = None
        self._pool.shutdown(wait=False)

    # ── Public API ────────────────────────────────────────────────────────────

    async def fetch_raw_option_chain(
        self,
        symbol: str,
        *,
        is_index: bool = True,
        expiry: str | None = None,
    ) -> dict[str, Any]:
        """Fetch the unparsed raw JSON dictionary for *symbol* from NSE."""
        symbol = symbol.upper().strip()
        self._breaker.raise_if_open(EXCHANGE, symbol)

        loop = asyncio.get_running_loop()
        try:
            return await asyncio.wait_for(
                loop.run_in_executor(
                    self._pool,
                    self._fetch_sync,
                    symbol,
                    is_index,
                    expiry,
                ),
                timeout=self._settings.async_timeout,
            )
        except asyncio.TimeoutError:
            self._breaker.record_failure(EXCHANGE, symbol)
            raise FetchTimeoutError(
                symbol=symbol,
                timeout_s=self._settings.async_timeout,
                exchange=EXCHANGE,
            )

    async def fetch_option_chain(
        self,
        symbol: str,
        *,
        is_index: bool = True,
        expiry: str | None = None,
    ) -> OptionChainResult:
        """Fetch the full option chain for *symbol*.

        Args:
            symbol:   NSE trading symbol (e.g. ``"NIFTY"``, ``"RELIANCE"``).
            is_index: ``True`` for index instruments (NIFTY, BANKNIFTY, …),
                      ``False`` for equity derivatives.
            expiry:   Specific expiry date string (``"27-Jun-2024"`` format).
                      Uses the nearest expiry when not provided.

        Returns:
            :class:`~indiaopt.models.option_chain.OptionChainResult`

        Raises:
            :class:`~indiaopt.exceptions.CircuitOpenError`: Circuit is open.
            :class:`~indiaopt.exceptions.FetchTimeoutError`: Exceeded async timeout.
            :class:`~indiaopt.exceptions.FetchError`: All retries exhausted.
        """
        symbol = symbol.upper().strip()
        self._breaker.raise_if_open(EXCHANGE, symbol)

        loop = asyncio.get_running_loop()
        try:
            raw = await asyncio.wait_for(
                loop.run_in_executor(
                    self._pool,
                    self._fetch_sync,
                    symbol,
                    is_index,
                    expiry,
                ),
                timeout=self._settings.async_timeout,
            )
        except asyncio.TimeoutError:
            self._breaker.record_failure(EXCHANGE, symbol)
            raise FetchTimeoutError(
                symbol=symbol,
                timeout_s=self._settings.async_timeout,
                exchange=EXCHANGE,
            )

        return nse_parser.parse_option_chain(raw, symbol)

    async def get_expiry_dates(self, symbol: str) -> list[str]:
        """Return available expiry dates for *symbol*.

        Args:
            symbol: NSE trading symbol.

        Returns:
            List of expiry date strings in ``"DD-Mon-YYYY"`` format.
        """
        symbol = symbol.upper().strip()
        loop = asyncio.get_running_loop()
        try:
            result = await asyncio.wait_for(
                loop.run_in_executor(
                    self._pool,
                    self._fetch_expiry_dates_sync,
                    symbol,
                ),
                timeout=self._settings.async_timeout,
            )
        except asyncio.TimeoutError:
            return []
        return result

    def reset_circuit(self, symbol: str) -> None:
        """Manually reset the circuit breaker for *symbol*.

        Use after manually resolving an issue (e.g. IP unblock, network fix).
        """
        self._breaker.reset(EXCHANGE, symbol)

    def circuit_stats(self) -> dict[str, dict[str, object]]:
        """Return a snapshot of all circuit breaker states for monitoring."""
        return self._breaker.stats()

    # ── Synchronous internals (run in thread pool) ────────────────────────────

    def _get_or_create_session(self) -> Any:
        """Return existing session or create a fresh one."""
        if self._session is None:
            self._session = create_nse_session(
                proxy_urls=self._settings.proxy_urls,
                warmup_timeout=self._settings.warmup_timeout,
            )
        return self._session

    def _fetch_sync(
        self,
        symbol: str,
        is_index: bool,
        expiry: str | None,
    ) -> dict[str, Any]:
        """Synchronous fetch with retry, circuit breaker, and session reuse.

        This runs in a thread pool worker.
        """
        session = self._get_or_create_session()
        type_str = "Indices" if is_index else "Equities"

        # Resolve expiry date if not provided
        if expiry is None:
            expiry = self._get_nearest_expiry_sync(session, symbol)

        url = (
            f"{NSE_BASE_URL}{NSE_OPTION_CHAIN_PATH}"
            f"?type={type_str}&symbol={symbol}&expiry={expiry}"
        )

        last_exc: Exception | None = None
        for attempt in range(self._retry_policy.max_retries):
            try:
                resp = session.get(url, timeout=self._settings.fetch_timeout)

                if resp.status_code == 200:
                    try:
                        data: dict[str, Any] = resp.json()
                    except Exception as json_exc:
                        # JSON parse failure — session cookies likely expired
                        logger.warning(
                            "JSON decode failed for %s (attempt %d): %s. "
                            "Refreshing session.",
                            symbol,
                            attempt + 1,
                            json_exc,
                        )
                        self._session = None
                        session = self._get_or_create_session()
                        raise ParseError(
                            f"Invalid JSON from NSE (likely cookie expiry): {json_exc}",
                            symbol=symbol,
                            exchange=EXCHANGE,
                            attempts=attempt + 1,
                            raw_preview=resp.text[:200] if hasattr(resp, "text") else "",
                        ) from json_exc

                    self._breaker.record_success(EXCHANGE, symbol)
                    logger.debug(
                        "Fetched %s/%s expiry=%s HTTP 200 (attempt %d)",
                        EXCHANGE,
                        symbol,
                        expiry,
                        attempt + 1,
                    )
                    return data

                if resp.status_code in RATE_LIMIT_STATUS_CODES:
                    retry_after: float | None = None
                    raw_ra = resp.headers.get("Retry-After") or resp.headers.get("retry-after")
                    if raw_ra:
                        try:
                            retry_after = float(raw_ra)
                        except ValueError:
                            pass

                    exc = RateLimitError(
                        f"HTTP {resp.status_code} from NSE for {symbol}",
                        symbol=symbol,
                        exchange=EXCHANGE,
                        status_code=resp.status_code,
                        attempts=attempt + 1,
                        retry_after=retry_after,
                    )
                    logger.warning(
                        "Rate-limited for %s (attempt %d/%d, HTTP %d). "
                        "Backing off.",
                        symbol,
                        attempt + 1,
                        self._retry_policy.max_retries,
                        resp.status_code,
                    )
                    last_exc = exc
                    if attempt < self._retry_policy.max_retries - 1:
                        self._retry_policy.sleep_for(attempt)
                    continue

                # Unexpected non-200 status — don't retry
                raise FetchError(
                    f"Unexpected HTTP {resp.status_code} for {symbol}",
                    symbol=symbol,
                    exchange=EXCHANGE,
                    status_code=resp.status_code,
                    attempts=attempt + 1,
                )

            except (FetchError, ParseError):
                raise  # propagate our own errors
            except OSError as exc:
                last_exc = NetworkError(
                    f"Network error for {symbol}: {exc}",
                    symbol=symbol,
                    exchange=EXCHANGE,
                    attempts=attempt + 1,
                    recovery="Check network connectivity and proxy configuration.",
                )
                logger.warning(
                    "Network error for %s (attempt %d/%d): %s",
                    symbol,
                    attempt + 1,
                    self._retry_policy.max_retries,
                    exc,
                )
                # Recreate session on network error
                self._session = None
                session = self._get_or_create_session()
                if attempt < self._retry_policy.max_retries - 1:
                    self._retry_policy.sleep_for(attempt)
            except Exception as exc:
                last_exc = exc
                logger.warning(
                    "Unexpected error for %s (attempt %d/%d): %s",
                    symbol,
                    attempt + 1,
                    self._retry_policy.max_retries,
                    exc,
                )
                if attempt < self._retry_policy.max_retries - 1:
                    self._retry_policy.sleep_for(attempt)

        self._breaker.record_failure(EXCHANGE, symbol)
        raise FetchError(
            f"All {self._retry_policy.max_retries} attempts failed for {symbol}",
            symbol=symbol,
            exchange=EXCHANGE,
            attempts=self._retry_policy.max_retries,
            recovery="Check NSE status and your network. Circuit breaker may open.",
        ) from last_exc

    def _get_nearest_expiry_sync(self, session: Any, symbol: str) -> str:
        """Fetch the nearest expiry date for *symbol* from NSE contract info.

        Falls back to today's date string if the API call fails.
        """
        from datetime import datetime

        fallback = datetime.now().strftime("%d-%b-%Y")
        url = f"{NSE_BASE_URL}{NSE_CONTRACT_INFO_PATH}?symbol={symbol}"
        try:
            resp = session.get(url, timeout=self._settings.fetch_timeout)
            if resp.status_code == 200:
                data = resp.json()
                dates: list[str] = data.get("expiryDates") or []
                if dates:
                    return dates[0]
        except Exception as exc:
            logger.debug("Could not fetch expiry dates for %s: %s. Using fallback.", symbol, exc)
        return fallback

    def _fetch_expiry_dates_sync(self, symbol: str) -> list[str]:
        """Synchronous helper to fetch expiry dates list."""
        session = self._get_or_create_session()
        url = f"{NSE_BASE_URL}{NSE_CONTRACT_INFO_PATH}?symbol={symbol}"
        try:
            resp = session.get(url, timeout=self._settings.fetch_timeout)
            if resp.status_code == 200:
                data = resp.json()
                return data.get("expiryDates") or []
        except Exception as exc:
            logger.debug("Could not fetch expiry dates for %s: %s", symbol, exc)
        return []


__all__ = ["NSEClient"]
