"""Thread-safe circuit breaker.

The circuit breaker prevents cascading failures by stopping fetch attempts
when a symbol repeatedly fails, allowing the exchange to recover.

States:

* **CLOSED** — normal operation; requests pass through.
* **OPEN** — failure threshold exceeded; requests are rejected immediately
  with :class:`~indiaopt.exceptions.CircuitOpenError`.
* **HALF-OPEN** — backoff period elapsed; one probe request is allowed
  through. Success → CLOSED, failure → OPEN (with reset backoff).

Each (exchange, symbol) pair has its own independent circuit state.

Example::

    breaker = CircuitBreaker(failure_threshold=5, backoff_seconds=300)

    # Inside your fetch loop:
    if breaker.is_open("NSE", "NIFTY"):
        raise CircuitOpenError(...)

    try:
        result = do_fetch()
        breaker.record_success("NSE", "NIFTY")
    except Exception:
        breaker.record_failure("NSE", "NIFTY")
        raise
"""

from __future__ import annotations

import logging
import threading
import time
from dataclasses import dataclass
from enum import Enum, auto
from typing import NamedTuple

from indiaopt.exceptions import CircuitOpenError

logger = logging.getLogger("indiaopt.circuit")


class CircuitState(Enum):
    """Possible states of a circuit."""

    CLOSED = auto()
    OPEN = auto()
    HALF_OPEN = auto()


class _CircuitKey(NamedTuple):
    exchange: str
    symbol: str


@dataclass
class _CircuitEntry:
    state: CircuitState = CircuitState.CLOSED
    failures: int = 0
    open_until: float = 0.0  # monotonic timestamp
    last_failure_time: float = 0.0

    def reset(self) -> None:
        self.state = CircuitState.CLOSED
        self.failures = 0
        self.open_until = 0.0


class CircuitBreaker:
    """Thread-safe circuit breaker with CLOSED / OPEN / HALF-OPEN states.

    Args:
        failure_threshold: Consecutive failures before opening the circuit.
        backoff_seconds:   How long (seconds) the circuit stays open before
                           moving to HALF-OPEN for a probe attempt.

    Usage::

        breaker = CircuitBreaker(failure_threshold=5, backoff_seconds=300)

        key = ("NSE", "NIFTY")
        if breaker.is_open(*key):
            raise CircuitOpenError("NIFTY", open_until=..., failures=...)

        try:
            result = do_fetch()
            breaker.record_success(*key)
        except Exception:
            breaker.record_failure(*key)
            raise
    """

    def __init__(
        self,
        failure_threshold: int = 5,
        backoff_seconds: float = 300.0,
    ) -> None:
        if failure_threshold < 1:
            raise ValueError("failure_threshold must be >= 1")
        if backoff_seconds <= 0:
            raise ValueError("backoff_seconds must be > 0")
        self._failure_threshold = failure_threshold
        self._backoff_seconds = backoff_seconds
        self._circuits: dict[_CircuitKey, _CircuitEntry] = {}
        self._lock = threading.Lock()

    # ── Internal helpers ──────────────────────────────────────────────────────

    def _get_entry(self, key: _CircuitKey) -> _CircuitEntry:
        """Return (or create) the entry for *key*. Must be called under lock."""
        if key not in self._circuits:
            self._circuits[key] = _CircuitEntry()
        return self._circuits[key]

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

    def is_open(self, exchange: str, symbol: str) -> bool:
        """Return ``True`` if the circuit is open for *(exchange, symbol)*.

        When the backoff period has elapsed, transitions the circuit to
        HALF-OPEN and returns ``False`` so one probe attempt can proceed.
        """
        key = _CircuitKey(exchange.upper(), symbol.upper())
        with self._lock:
            entry = self._get_entry(key)
            if entry.state == CircuitState.CLOSED:
                return False
            if entry.state == CircuitState.OPEN:
                now = time.monotonic()
                if now >= entry.open_until:
                    # Backoff elapsed — allow one probe attempt
                    entry.state = CircuitState.HALF_OPEN
                    logger.info(
                        "Circuit HALF-OPEN for %s/%s — allowing probe attempt.",
                        exchange,
                        symbol,
                    )
                    return False
                return True
            # HALF_OPEN — let the probe through
            return False

    def get_state(self, exchange: str, symbol: str) -> CircuitState:
        """Return the current :class:`CircuitState` for *(exchange, symbol)*."""
        key = _CircuitKey(exchange.upper(), symbol.upper())
        with self._lock:
            return self._get_entry(key).state

    def record_success(self, exchange: str, symbol: str) -> None:
        """Record a successful fetch for *(exchange, symbol)*.

        Resets the circuit to CLOSED regardless of prior state.
        """
        key = _CircuitKey(exchange.upper(), symbol.upper())
        with self._lock:
            entry = self._get_entry(key)
            if entry.failures > 0 or entry.state != CircuitState.CLOSED:
                logger.info(
                    "Circuit CLOSED for %s/%s after successful fetch "
                    "(had %d failures).",
                    exchange,
                    symbol,
                    entry.failures,
                )
            entry.reset()

    def record_failure(self, exchange: str, symbol: str) -> None:
        """Record a fetch failure for *(exchange, symbol)*.

        Opens the circuit when failures reach the threshold.
        """
        key = _CircuitKey(exchange.upper(), symbol.upper())
        with self._lock:
            entry = self._get_entry(key)
            entry.failures += 1
            entry.last_failure_time = time.monotonic()
            if entry.state == CircuitState.HALF_OPEN:
                # Probe failed — reopen
                entry.state = CircuitState.OPEN
                entry.open_until = time.monotonic() + self._backoff_seconds
                logger.warning(
                    "Circuit re-OPENED for %s/%s after probe failure. "
                    "Backing off for %.0fs.",
                    exchange,
                    symbol,
                    self._backoff_seconds,
                )
            elif entry.failures >= self._failure_threshold:
                entry.state = CircuitState.OPEN
                entry.open_until = time.monotonic() + self._backoff_seconds
                logger.warning(
                    "Circuit OPENED for %s/%s after %d failures. "
                    "Backing off for %.0fs.",
                    exchange,
                    symbol,
                    entry.failures,
                    self._backoff_seconds,
                )

    def reset(self, exchange: str, symbol: str) -> None:
        """Manually reset the circuit for *(exchange, symbol)* to CLOSED.

        Useful in tests or operator-driven recovery scenarios.
        """
        key = _CircuitKey(exchange.upper(), symbol.upper())
        with self._lock:
            if key in self._circuits:
                self._circuits[key].reset()
        logger.info("Circuit manually reset for %s/%s.", exchange, symbol)

    def reset_all(self) -> None:
        """Reset ALL circuits to CLOSED. Use with caution in production."""
        with self._lock:
            for entry in self._circuits.values():
                entry.reset()
        logger.info("All circuits manually reset.")

    def raise_if_open(self, exchange: str, symbol: str) -> None:
        """Raise :class:`~indiaopt.exceptions.CircuitOpenError` if circuit is open.

        Convenience wrapper for the common check-and-raise pattern::

            breaker.raise_if_open("NSE", "NIFTY")  # raises or passes through
        """
        key = _CircuitKey(exchange.upper(), symbol.upper())
        with self._lock:
            entry = self._get_entry(key)
            if entry.state == CircuitState.OPEN:
                now = time.monotonic()
                if now < entry.open_until:
                    raise CircuitOpenError(
                        symbol=symbol,
                        open_until=entry.open_until,
                        failures=entry.failures,
                        exchange=exchange,
                    )
                # Backoff elapsed — transition to HALF-OPEN silently
                entry.state = CircuitState.HALF_OPEN
                logger.info(
                    "Circuit HALF-OPEN for %s/%s — allowing probe attempt.",
                    exchange,
                    symbol,
                )

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

        Returns:
            Dict keyed by ``"EXCHANGE/SYMBOL"`` with state info.
        """
        with self._lock:
            return {
                f"{key.exchange}/{key.symbol}": {
                    "state": entry.state.name,
                    "failures": entry.failures,
                    "open_until": entry.open_until,
                }
                for key, entry in self._circuits.items()
            }


# Module-level shared instance — clients can share this or create their own
_default_breaker: CircuitBreaker | None = None
_default_breaker_lock = threading.Lock()


def get_default_breaker(
    failure_threshold: int = 5,
    backoff_seconds: float = 300.0,
) -> CircuitBreaker:
    """Return the module-level shared :class:`CircuitBreaker`.

    Creates it on first call with the given parameters. Subsequent calls
    return the same instance regardless of parameters.
    """
    global _default_breaker
    with _default_breaker_lock:
        if _default_breaker is None:
            _default_breaker = CircuitBreaker(
                failure_threshold=failure_threshold,
                backoff_seconds=backoff_seconds,
            )
    return _default_breaker


__all__ = ["CircuitBreaker", "CircuitState", "get_default_breaker"]
