"""HTTP layer — retry policy with exponential backoff and jitter.

Key design decisions vs. original code:
- ``_MAX_RETRIES = 1000`` → configurable, default 5, max 10.
- ``2**attempt`` unbounded sleep → ``min(base**attempt + jitter, cap)`` capped at 30s.
- Session is NOT recreated on every retry — only recreated on explicit session error.
"""

from __future__ import annotations

import logging
import random
import time
from collections.abc import Callable
from dataclasses import dataclass
from typing import TypeVar

from indiaopt.constants import (
    DEFAULT_BACKOFF_BASE,
    DEFAULT_BACKOFF_CAP_S,
    DEFAULT_BACKOFF_JITTER,
    DEFAULT_MAX_RETRIES,
    RATE_LIMIT_STATUS_CODES,
)

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

T = TypeVar("T")


@dataclass(frozen=True)
class RetryPolicy:
    """Immutable retry policy configuration.

    Attributes:
        max_retries:    Maximum number of attempts (including the first).
        backoff_base:   Base for exponential backoff formula.
        backoff_cap:    Maximum sleep duration in seconds.
        jitter:         Maximum random jitter added to each sleep.
        retry_statuses: HTTP status codes that trigger a retry.

    Backoff formula::

        sleep = min(base ** attempt + uniform(0, jitter), cap)
    """

    max_retries: int = DEFAULT_MAX_RETRIES
    backoff_base: float = DEFAULT_BACKOFF_BASE
    backoff_cap: float = DEFAULT_BACKOFF_CAP_S
    jitter: float = DEFAULT_BACKOFF_JITTER
    retry_statuses: frozenset[int] = RATE_LIMIT_STATUS_CODES

    def sleep_for(self, attempt: int) -> float:
        """Compute and sleep for the backoff duration on *attempt* (0-indexed).

        Returns:
            Actual sleep duration in seconds.
        """
        duration = min(
            self.backoff_base**attempt + random.uniform(0.0, self.jitter),
            self.backoff_cap,
        )
        logger.debug("Backoff sleep %.2fs (attempt %d).", duration, attempt + 1)
        time.sleep(duration)
        return duration

    def compute_sleep(self, attempt: int) -> float:
        """Return the backoff duration for *attempt* without sleeping.

        Useful for logging and testing.
        """
        return min(
            self.backoff_base**attempt + random.uniform(0.0, self.jitter),
            self.backoff_cap,
        )

    def should_retry_status(self, status_code: int) -> bool:
        """Return ``True`` if *status_code* warrants a retry."""
        return status_code in self.retry_statuses

    @classmethod
    def from_settings(cls, settings: object) -> RetryPolicy:
        """Construct a :class:`RetryPolicy` from a :class:`~indiaopt.config.settings.Settings` object."""
        return cls(
            max_retries=getattr(settings, "max_retries", DEFAULT_MAX_RETRIES),
            backoff_base=getattr(settings, "backoff_base", DEFAULT_BACKOFF_BASE),
            backoff_cap=getattr(settings, "backoff_cap_seconds", DEFAULT_BACKOFF_CAP_S),
            jitter=getattr(settings, "backoff_jitter", DEFAULT_BACKOFF_JITTER),
        )


def with_retry(
    fn: Callable[[], T],
    policy: RetryPolicy,
    *,
    label: str = "operation",
    on_retry: Callable[[int, Exception], None] | None = None,
) -> T:
    """Execute *fn* with retry logic defined by *policy*.

    Args:
        fn:       Zero-argument callable to execute.
        policy:   :class:`RetryPolicy` governing retry behaviour.
        label:    Human-readable name for logging.
        on_retry: Optional callback invoked with ``(attempt, exc)`` before sleep.

    Returns:
        The return value of *fn* on success.

    Raises:
        The last exception raised by *fn* after all attempts are exhausted.

    Example::

        result = with_retry(
            lambda: session.get(url, timeout=15),
            RetryPolicy(max_retries=5),
            label=f"fetch {symbol}",
        )
    """
    last_exc: Exception | None = None
    for attempt in range(policy.max_retries):
        try:
            return fn()
        except Exception as exc:
            last_exc = exc
            if on_retry is not None:
                on_retry(attempt, exc)
            logger.warning(
                "Attempt %d/%d for '%s' failed: %s",
                attempt + 1,
                policy.max_retries,
                label,
                exc,
            )
            if attempt < policy.max_retries - 1:
                policy.sleep_for(attempt)

    # Should never be None if max_retries >= 1, but satisfies type checker
    raise last_exc or RuntimeError(f"All {policy.max_retries} attempts for '{label}' failed.")


__all__ = ["RetryPolicy", "with_retry"]
