"""
Circuit Breaker Pattern Implementation.

الگوی Circuit Breaker برای محافظت از سرویس‌ها در مقابل خطاهای زنجیره‌ای.

States:
    CLOSED  → عملکرد عادی — درخواست‌ها عبور می‌کنند
    OPEN    → خطا تشخیص داده شده — درخواست‌ها بلافاصله رد می‌شوند
    HALF_OPEN → دوره آزمایشی — تعداد محدودی درخواست عبور می‌کنند

Usage:
    breaker = CircuitBreaker(
        name='external-api',
        failure_threshold=5,
        recovery_timeout=60,
    )

    result = breaker.call(my_function, arg1, arg2)

    # With decorator
    @circuit_breaker(name='payment-gateway', failure_threshold=3)
    def process_payment(amount):
        ...

    # With fallback
    breaker = CircuitBreaker(
        name='search',
        failure_threshold=3,
        fallback=lambda *a, **kw: [],
    )
"""
import enum
import functools
import logging
import threading
import time
from dataclasses import dataclass, field
from typing import Any, Callable, Optional

logger = logging.getLogger(__name__)


# =============================================================================
# Exceptions
# =============================================================================

class CircuitBreakerError(Exception):
    """Base exception for circuit breaker errors."""
    pass


class CircuitOpenError(CircuitBreakerError):
    """Raised when circuit breaker is open and call is rejected."""

    def __init__(self, breaker_name: str, remaining_seconds: float = 0):
        self.breaker_name = breaker_name
        self.remaining_seconds = remaining_seconds
        super().__init__(
            f'Circuit breaker "{breaker_name}" is OPEN. '
            f'Retry after {remaining_seconds:.1f}s.'
        )


# =============================================================================
# Circuit Breaker State
# =============================================================================

class CircuitState(enum.Enum):
    CLOSED = 'closed'
    OPEN = 'open'
    HALF_OPEN = 'half_open'


@dataclass
class CircuitBreakerMetrics:
    """Metrics and stats for a circuit breaker instance."""
    total_calls: int = 0
    successful_calls: int = 0
    failed_calls: int = 0
    rejected_calls: int = 0
    consecutive_failures: int = 0
    consecutive_successes: int = 0
    last_failure_time: Optional[float] = None
    last_success_time: Optional[float] = None
    last_state_change_time: Optional[float] = None
    state_change_count: int = 0
    last_error: Optional[str] = None
    last_error_type: Optional[str] = None

    def reset_counters(self):
        self.consecutive_failures = 0
        self.consecutive_successes = 0

    def record_success(self):
        self.total_calls += 1
        self.successful_calls += 1
        self.consecutive_successes += 1
        self.consecutive_failures = 0
        self.last_success_time = time.monotonic()

    def record_failure(self, error: Exception):
        self.total_calls += 1
        self.failed_calls += 1
        self.consecutive_failures += 1
        self.consecutive_successes = 0
        self.last_failure_time = time.monotonic()
        self.last_error = str(error)
        self.last_error_type = type(error).__name__

    def record_rejection(self):
        self.total_calls += 1
        self.rejected_calls += 1

    def record_state_change(self):
        self.last_state_change_time = time.monotonic()
        self.state_change_count += 1

    @property
    def failure_rate(self) -> float:
        actual_calls = self.successful_calls + self.failed_calls
        if actual_calls == 0:
            return 0.0
        return self.failed_calls / actual_calls

    def to_dict(self) -> dict:
        return {
            'total_calls': self.total_calls,
            'successful_calls': self.successful_calls,
            'failed_calls': self.failed_calls,
            'rejected_calls': self.rejected_calls,
            'failure_rate': round(self.failure_rate, 4),
            'consecutive_failures': self.consecutive_failures,
            'consecutive_successes': self.consecutive_successes,
            'last_error': self.last_error,
            'last_error_type': self.last_error_type,
            'state_change_count': self.state_change_count,
        }


# =============================================================================
# Circuit Breaker
# =============================================================================

class CircuitBreaker:
    """
    Circuit Breaker implementation with three states.

    Args:
        name: Unique identifier for this breaker.
        failure_threshold: Number of consecutive failures before opening.
        recovery_timeout: Seconds to wait before transitioning to HALF_OPEN.
        success_threshold: Consecutive successes in HALF_OPEN to close.
        half_open_max_calls: Max concurrent calls allowed in HALF_OPEN.
        fallback: Optional fallback function when circuit is open.
        excluded_exceptions: Exception types that should NOT trip the breaker.
        on_open: Callback when circuit opens.
        on_close: Callback when circuit closes.
        on_half_open: Callback when circuit enters half-open.
    """

    # Global registry of all circuit breakers
    _registry: dict[str, 'CircuitBreaker'] = {}
    _registry_lock = threading.Lock()

    def __init__(
        self,
        name: str,
        failure_threshold: int = 5,
        recovery_timeout: float = 60.0,
        success_threshold: int = 3,
        half_open_max_calls: int = 1,
        fallback: Optional[Callable] = None,
        excluded_exceptions: Optional[tuple[type[Exception], ...]] = None,
        on_open: Optional[Callable] = None,
        on_close: Optional[Callable] = None,
        on_half_open: Optional[Callable] = None,
    ):
        self.name = name
        self.failure_threshold = failure_threshold
        self.recovery_timeout = recovery_timeout
        self.success_threshold = success_threshold
        self.half_open_max_calls = half_open_max_calls
        self.fallback = fallback
        self.excluded_exceptions = excluded_exceptions or ()
        self.on_open = on_open
        self.on_close = on_close
        self.on_half_open = on_half_open

        self._state = CircuitState.CLOSED
        self._lock = threading.RLock()
        self._metrics = CircuitBreakerMetrics()
        self._opened_at: Optional[float] = None
        self._half_open_calls = 0

        # Register in global registry
        with self._registry_lock:
            self._registry[name] = self

    # -------------------------------------------------------------------------
    # Properties
    # -------------------------------------------------------------------------

    @property
    def state(self) -> CircuitState:
        with self._lock:
            if self._state == CircuitState.OPEN:
                if self._should_transition_to_half_open():
                    self._transition_to(CircuitState.HALF_OPEN)
            return self._state

    @property
    def metrics(self) -> CircuitBreakerMetrics:
        return self._metrics

    @property
    def is_closed(self) -> bool:
        return self.state == CircuitState.CLOSED

    @property
    def is_open(self) -> bool:
        return self.state == CircuitState.OPEN

    # -------------------------------------------------------------------------
    # Core Methods
    # -------------------------------------------------------------------------

    def call(self, func: Callable, *args, **kwargs) -> Any:
        """
        Execute a function through the circuit breaker.

        Raises CircuitOpenError if breaker is open and no fallback is set.
        """
        with self._lock:
            current_state = self.state

            if current_state == CircuitState.OPEN:
                self._metrics.record_rejection()
                return self._handle_open(func, *args, **kwargs)

            if current_state == CircuitState.HALF_OPEN:
                if self._half_open_calls >= self.half_open_max_calls:
                    self._metrics.record_rejection()
                    return self._handle_open(func, *args, **kwargs)
                self._half_open_calls += 1

        # Execute outside lock to avoid blocking
        try:
            result = func(*args, **kwargs)
            self._on_success()
            return result
        except Exception as e:
            if isinstance(e, self.excluded_exceptions):
                # Excluded exceptions don't count as failures
                self._metrics.record_success()
                raise
            self._on_failure(e)
            raise

    def _handle_open(self, func: Callable, *args, **kwargs) -> Any:
        """Handle a call when breaker is open."""
        if self.fallback:
            logger.warning(
                'Circuit breaker "%s" is OPEN, using fallback.',
                self.name,
            )
            return self.fallback(*args, **kwargs)

        remaining = self._get_remaining_timeout()
        raise CircuitOpenError(self.name, remaining)

    def _on_success(self):
        """Record a successful call."""
        with self._lock:
            self._metrics.record_success()

            if self._state == CircuitState.HALF_OPEN:
                if self._metrics.consecutive_successes >= self.success_threshold:
                    self._transition_to(CircuitState.CLOSED)

    def _on_failure(self, error: Exception):
        """Record a failed call."""
        with self._lock:
            self._metrics.record_failure(error)

            logger.warning(
                'Circuit breaker "%s" recorded failure (%d/%d): %s',
                self.name,
                self._metrics.consecutive_failures,
                self.failure_threshold,
                str(error),
            )

            if self._state == CircuitState.HALF_OPEN:
                # Any failure in half-open immediately re-opens
                self._transition_to(CircuitState.OPEN)
            elif self._state == CircuitState.CLOSED:
                if self._metrics.consecutive_failures >= self.failure_threshold:
                    self._transition_to(CircuitState.OPEN)

    # -------------------------------------------------------------------------
    # State Transitions
    # -------------------------------------------------------------------------

    def _transition_to(self, new_state: CircuitState):
        """Transition to a new state."""
        old_state = self._state
        self._state = new_state
        self._metrics.record_state_change()
        self._metrics.reset_counters()

        logger.info(
            'Circuit breaker "%s" transitioned: %s → %s',
            self.name,
            old_state.value,
            new_state.value,
        )

        if new_state == CircuitState.OPEN:
            self._opened_at = time.monotonic()
            self._half_open_calls = 0
            if self.on_open:
                self._safe_callback(self.on_open, self)

        elif new_state == CircuitState.HALF_OPEN:
            self._half_open_calls = 0
            if self.on_half_open:
                self._safe_callback(self.on_half_open, self)

        elif new_state == CircuitState.CLOSED:
            self._opened_at = None
            self._half_open_calls = 0
            if self.on_close:
                self._safe_callback(self.on_close, self)

    def _should_transition_to_half_open(self) -> bool:
        if self._opened_at is None:
            return False
        elapsed = time.monotonic() - self._opened_at
        return elapsed >= self.recovery_timeout

    def _get_remaining_timeout(self) -> float:
        if self._opened_at is None:
            return 0.0
        elapsed = time.monotonic() - self._opened_at
        return max(0.0, self.recovery_timeout - elapsed)

    @staticmethod
    def _safe_callback(callback: Callable, *args):
        try:
            callback(*args)
        except Exception as e:
            logger.error('Circuit breaker callback error: %s', e)

    # -------------------------------------------------------------------------
    # Manual Controls
    # -------------------------------------------------------------------------

    def reset(self):
        """Manually reset the circuit breaker to CLOSED state."""
        with self._lock:
            self._transition_to(CircuitState.CLOSED)
            self._metrics = CircuitBreakerMetrics()
            logger.info('Circuit breaker "%s" manually reset.', self.name)

    def force_open(self):
        """Manually force the circuit breaker to OPEN state."""
        with self._lock:
            self._transition_to(CircuitState.OPEN)
            logger.info('Circuit breaker "%s" manually opened.', self.name)

    def get_status(self) -> dict:
        """Get the full status of this circuit breaker."""
        return {
            'name': self.name,
            'state': self.state.value,
            'failure_threshold': self.failure_threshold,
            'recovery_timeout': self.recovery_timeout,
            'success_threshold': self.success_threshold,
            'metrics': self._metrics.to_dict(),
        }

    # -------------------------------------------------------------------------
    # Registry Class Methods
    # -------------------------------------------------------------------------

    @classmethod
    def get_breaker(cls, name: str) -> Optional['CircuitBreaker']:
        """Get a circuit breaker by name from the global registry."""
        return cls._registry.get(name)

    @classmethod
    def get_all_breakers(cls) -> dict[str, 'CircuitBreaker']:
        """Get all registered circuit breakers."""
        return dict(cls._registry)

    @classmethod
    def get_all_statuses(cls) -> list[dict]:
        """Get status of all circuit breakers."""
        return [breaker.get_status() for breaker in cls._registry.values()]

    @classmethod
    def reset_all(cls):
        """Reset all circuit breakers."""
        for breaker in cls._registry.values():
            breaker.reset()


# =============================================================================
# Decorator
# =============================================================================

def circuit_breaker(
    name: str,
    failure_threshold: int = 5,
    recovery_timeout: float = 60.0,
    success_threshold: int = 3,
    fallback: Optional[Callable] = None,
    excluded_exceptions: Optional[tuple[type[Exception], ...]] = None,
):
    """
    Decorator to wrap a function with a circuit breaker.

    Usage:
        @circuit_breaker(name='external-api', failure_threshold=3)
        def call_external_api(data):
            ...
    """
    breaker = CircuitBreaker.get_breaker(name)
    if breaker is None:
        breaker = CircuitBreaker(
            name=name,
            failure_threshold=failure_threshold,
            recovery_timeout=recovery_timeout,
            success_threshold=success_threshold,
            fallback=fallback,
            excluded_exceptions=excluded_exceptions,
        )

    def decorator(func: Callable):
        @functools.wraps(func)
        def wrapper(*args, **kwargs):
            return breaker.call(func, *args, **kwargs)

        wrapper.circuit_breaker = breaker
        return wrapper

    return decorator


# =============================================================================
# Django Integration Mixin
# =============================================================================

class CircuitBreakerMixin:
    """
    Mixin for Django views / DRF views to use circuit breaker on external calls.

    Usage:
        class PaymentView(CircuitBreakerMixin, APIView):
            circuit_breaker_name = 'payment-gateway'
            circuit_breaker_failure_threshold = 3
            circuit_breaker_recovery_timeout = 30
            circuit_breaker_fallback = None

            def post(self, request):
                result = self.call_with_breaker(
                    self.process_payment, request.data
                )
                return Response(result)

            def process_payment(self, data):
                # call external API
                ...
    """

    circuit_breaker_name: str = 'default'
    circuit_breaker_failure_threshold: int = 5
    circuit_breaker_recovery_timeout: float = 60.0
    circuit_breaker_success_threshold: int = 3
    circuit_breaker_fallback: Optional[Callable] = None

    def get_circuit_breaker(self) -> CircuitBreaker:
        name = self.circuit_breaker_name
        breaker = CircuitBreaker.get_breaker(name)
        if breaker is None:
            breaker = CircuitBreaker(
                name=name,
                failure_threshold=self.circuit_breaker_failure_threshold,
                recovery_timeout=self.circuit_breaker_recovery_timeout,
                success_threshold=self.circuit_breaker_success_threshold,
                fallback=self.circuit_breaker_fallback,
            )
        return breaker

    def call_with_breaker(self, func: Callable, *args, **kwargs) -> Any:
        breaker = self.get_circuit_breaker()
        return breaker.call(func, *args, **kwargs)
