#!/usr/bin/env python3
"""
Ejemplo de un cliente de IA resiliente en un solo archivo.

Incluye:
- timeout
- retry con exponential backoff + jitter
- circuit breaker (CLOSED / OPEN / HALF_OPEN)
- fallback
- métricas
- logs
- ejemplo ejecutable sin dependencias externas
- adaptador opcional para OpenAI

Uso rápido:
    python resilient_ai_example.py

Uso con OpenAI:
    pip install openai
    set OPENAI_API_KEY=...        # Windows CMD
    $env:OPENAI_API_KEY="..."     # PowerShell
    export OPENAI_API_KEY="..."   # Linux/macOS

    python resilient_ai_example.py --openai "Explica circuit breaker en 3 frases"
"""

from __future__ import annotations

import argparse
import logging
import os
import random
import time
from dataclasses import dataclass
from enum import StrEnum
from typing import Protocol

# ---------------------------------------------------------------------------
# Logging
# ---------------------------------------------------------------------------

logging.basicConfig(
    level=logging.INFO,
    format="%(asctime)s | %(levelname)s | %(name)s | %(message)s",
)
logger = logging.getLogger("resilient-ai")


# ---------------------------------------------------------------------------
# Errores
# ---------------------------------------------------------------------------

class AIError(Exception):
    """Error base del proveedor de IA."""


class RetryableAIError(AIError):
    """Error temporal: vale la pena reintentar."""


class NonRetryableAIError(AIError):
    """Error permanente: no tiene sentido reintentar."""


class CircuitOpenError(AIError):
    """El circuit breaker está abierto."""


# ---------------------------------------------------------------------------
# Contrato del proveedor
# ---------------------------------------------------------------------------

class AIProvider(Protocol):
    def generate(self, prompt: str, timeout: float) -> str:
        """
        Debe generar una respuesta y respetar `timeout`.

        Debe lanzar:
        - RetryableAIError para errores temporales
        - NonRetryableAIError para errores permanentes
        """
        ...


# ---------------------------------------------------------------------------
# Métricas simples
# ---------------------------------------------------------------------------

@dataclass
class Metrics:
    requests: int = 0
    successful_requests: int = 0
    failed_requests: int = 0
    retries: int = 0
    fallbacks: int = 0
    circuit_rejections: int = 0
    provider_failures: int = 0

    def log(self) -> None:
        logger.info(
            "metrics requests=%d success=%d failed=%d retries=%d "
            "fallbacks=%d circuit_rejections=%d provider_failures=%d",
            self.requests,
            self.successful_requests,
            self.failed_requests,
            self.retries,
            self.fallbacks,
            self.circuit_rejections,
            self.provider_failures,
        )


# ---------------------------------------------------------------------------
# Circuit breaker
# ---------------------------------------------------------------------------

class CircuitState(StrEnum):
    CLOSED = "CLOSED"
    OPEN = "OPEN"
    HALF_OPEN = "HALF_OPEN"


class CircuitBreaker:
    """
    CLOSED:
        Las llamadas pasan normalmente.

    OPEN:
        Las llamadas al proveedor son rechazadas temporalmente.

    HALF_OPEN:
        Se permite una llamada de prueba.
        Si funciona -> CLOSED.
        Si falla -> OPEN.
    """

    def __init__(
        self,
        failure_threshold: int = 3,
        recovery_timeout: float = 15.0,
    ) -> None:
        self.failure_threshold = failure_threshold
        self.recovery_timeout = recovery_timeout

        self.state = CircuitState.CLOSED
        self.failure_count = 0
        self.opened_at: float | None = None

    def allow_request(self) -> bool:
        if self.state == CircuitState.CLOSED:
            return True

        if self.state == CircuitState.OPEN:
            assert self.opened_at is not None

            elapsed = time.monotonic() - self.opened_at
            if elapsed >= self.recovery_timeout:
                self.state = CircuitState.HALF_OPEN
                logger.warning("Circuit breaker: OPEN -> HALF_OPEN")
                return True

            return False

        # HALF_OPEN: en este ejemplo simple se permite la llamada de prueba.
        return True

    def record_success(self) -> None:
        previous = self.state
        self.failure_count = 0
        self.opened_at = None
        self.state = CircuitState.CLOSED

        if previous != CircuitState.CLOSED:
            logger.info("Circuit breaker: %s -> CLOSED", previous)

    def record_failure(self) -> None:
        self.failure_count += 1

        if (
            self.state == CircuitState.HALF_OPEN
            or self.failure_count >= self.failure_threshold
        ):
            self.state = CircuitState.OPEN
            self.opened_at = time.monotonic()
            logger.error(
                "Circuit breaker OPEN: failure_count=%d",
                self.failure_count,
            )


# ---------------------------------------------------------------------------
# Cliente resiliente
# ---------------------------------------------------------------------------

class ResilientAIClient:
    def __init__(
        self,
        provider: AIProvider,
        *,
        timeout: float = 5.0,
        max_attempts: int = 3,
        initial_backoff: float = 0.5,
        max_backoff: float = 8.0,
        jitter: float = 0.25,
        circuit_breaker: CircuitBreaker | None = None,
    ) -> None:
        self.provider = provider
        self.timeout = timeout
        self.max_attempts = max_attempts
        self.initial_backoff = initial_backoff
        self.max_backoff = max_backoff
        self.jitter = jitter
        self.circuit_breaker = circuit_breaker or CircuitBreaker()
        self.metrics = Metrics()

    def generate(self, prompt: str) -> str:
        """
        Flujo:

            request
               ↓
            timeout
               ↓
            retry + exponential backoff
               ↓
            circuit breaker
               ↓
            fallback
               ↓
            métricas/logs
        """

        self.metrics.requests += 1

        if not self.circuit_breaker.allow_request():
            self.metrics.circuit_rejections += 1
            logger.warning("Request rechazado: circuit breaker OPEN")

            result = self._fallback(
                prompt,
                CircuitOpenError("Circuit breaker is open"),
            )
            self.metrics.failed_requests += 1
            self.metrics.log()
            return result

        last_error: Exception | None = None

        for attempt in range(1, self.max_attempts + 1):
            try:
                logger.info(
                    "Llamando al proveedor attempt=%d/%d timeout=%.1fs",
                    attempt,
                    self.max_attempts,
                    self.timeout,
                )

                response = self.provider.generate(
                    prompt,
                    timeout=self.timeout,
                )

                self.circuit_breaker.record_success()
                self.metrics.successful_requests += 1
                self.metrics.log()
                return response

            except NonRetryableAIError as exc:
                last_error = exc
                self.metrics.provider_failures += 1
                logger.error("Error no reintentable: %s", exc)

                # Un 4xx funcional, por ejemplo, no suele significar que
                # el servicio entero esté caído, así que no abrimos el circuito.
                break

            except (RetryableAIError, TimeoutError, ConnectionError, OSError) as exc:
                last_error = exc
                self.metrics.provider_failures += 1
                self.circuit_breaker.record_failure()

                logger.warning(
                    "Fallo temporal attempt=%d/%d error=%r",
                    attempt,
                    self.max_attempts,
                    exc,
                )

                # Si el circuit breaker se abrió, no seguimos golpeando el servicio.
                if self.circuit_breaker.state == CircuitState.OPEN:
                    break

                if attempt < self.max_attempts:
                    self.metrics.retries += 1
                    delay = self._backoff_delay(attempt)
                    logger.info("Retry en %.2fs", delay)
                    time.sleep(delay)

            except Exception as exc:
                # Error inesperado: se registra y se usa fallback,
                # evitando un retry ciego de algo posiblemente no temporal.
                last_error = exc
                self.metrics.provider_failures += 1
                logger.exception("Error inesperado del proveedor")
                break

        self.metrics.failed_requests += 1

        result = self._fallback(
            prompt,
            last_error or AIError("Unknown provider error"),
        )
        self.metrics.log()
        return result

    def _backoff_delay(self, attempt: int) -> float:
        """
        Exponential backoff:

            0.5s, 1s, 2s, 4s...

        más un jitter aleatorio para evitar thundering herd.
        """
        base = min(
            self.max_backoff,
            self.initial_backoff * (2 ** (attempt - 1)),
        )

        jitter_amount = random.uniform(0.0, self.jitter)
        return base + jitter_amount

    def _fallback(self, prompt: str, error: Exception) -> str:
        self.metrics.fallbacks += 1

        logger.error(
            "Usando fallback. prompt=%r error=%r",
            prompt[:80],
            error,
        )

        # En producción el fallback podría ser:
        # - otro modelo/proveedor
        # - una respuesta cacheada
        # - una respuesta determinística
        # - una cola para reprocesamiento posterior
        return (
            "[FALLBACK] El servicio de IA no está disponible temporalmente. "
            "La solicitud fue manejada sin tumbar la aplicación."
        )


# ---------------------------------------------------------------------------
# Proveedor de demostración
# ---------------------------------------------------------------------------

class DemoProvider:
    """
    Simula un proveedor inestable para que el archivo pueda ejecutarse
    sin instalar nada ni usar una API real.
    """

    def __init__(self, failure_probability: float = 0.45) -> None:
        self.failure_probability = failure_probability

    def generate(self, prompt: str, timeout: float) -> str:
        # Simulamos una pequeña latencia.
        simulated_latency = random.uniform(0.05, 0.25)

        if simulated_latency > timeout:
            raise TimeoutError(f"Timeout después de {timeout}s")

        time.sleep(simulated_latency)

        if random.random() < self.failure_probability:
            raise RetryableAIError("HTTP 503 / servicio temporalmente no disponible")

        return f"Respuesta IA: {prompt}"


# ---------------------------------------------------------------------------
# Adaptador opcional para OpenAI
# ---------------------------------------------------------------------------

class OpenAIProvider:
    """
    Requiere:
        pip install openai
        OPENAI_API_KEY configurada

    Los retries del SDK se desactivan para que la política de retries
    quede centralizada en ResilientAIClient.
    """

    def __init__(self, model: str = "gpt-5.6") -> None:
        try:
            from openai import OpenAI
        except ImportError as exc:
            raise RuntimeError(
                "Falta el paquete 'openai'. Instálalo con: pip install openai"
            ) from exc

        if not os.getenv("OPENAI_API_KEY"):
            raise RuntimeError("OPENAI_API_KEY no está configurada")

        self._OpenAI = OpenAI
        self.model = model

    def generate(self, prompt: str, timeout: float) -> str:
        # Creamos el cliente con timeout y retries internos desactivados.
        client = self._OpenAI(
            timeout=timeout,
            max_retries=0,
        )

        try:
            response = client.responses.create(
                model=self.model,
                input=prompt,
            )
            return response.output_text

        except Exception as exc:
            # Clasificación básica sin acoplar demasiado este ejemplo
            # a las clases internas/versiones del SDK.
            status_code = getattr(exc, "status_code", None)

            if status_code in {408, 409, 429, 500, 502, 503, 504}:
                raise RetryableAIError(
                    f"OpenAI error temporal HTTP {status_code}"
                ) from exc

            if isinstance(status_code, int) and 400 <= status_code < 500:
                raise NonRetryableAIError(
                    f"OpenAI error no reintentable HTTP {status_code}"
                ) from exc

            # Errores de red/timeout/transport desconocidos se consideran
            # temporalmente reintentables.
            raise RetryableAIError(str(exc)) from exc


# ---------------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------------

def build_client(use_openai: bool) -> ResilientAIClient:
    provider: AIProvider = OpenAIProvider() if use_openai else DemoProvider()

    return ResilientAIClient(
        provider=provider,
        timeout=5.0,
        max_attempts=3,
        initial_backoff=0.5,
        max_backoff=4.0,
        jitter=0.3,
        circuit_breaker=CircuitBreaker(
            failure_threshold=3,
            recovery_timeout=10.0,
        ),
    )


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument(
        "prompt",
        nargs="?",
        default="Explica qué es ingeniería de confiabilidad.",
    )
    parser.add_argument(
        "--openai",
        action="store_true",
        help="Usar OpenAI en lugar del proveedor simulado.",
    )
    args = parser.parse_args()

    client = build_client(use_openai=args.openai)

    response = client.generate(args.prompt)
    print()
    print(response)


if __name__ == "__main__":
    main()
