Módulo 6: Code Quality Patterns para AI

5. Dependency Injection para LLM Providers

Descripción

Dependency Injection (DI) es el pattern que desacopla tu business logic de cualquier implementación específica del LLM. Sin DI, cambiar de OpenAI a Anthropic significa reescribir código de negocio. Con DI, es crear una nueva implementación del Protocol y cambiar una línea en la configuración. Esta cápsula implementa DI completo: el Protocol, las implementaciones (OpenAI, Mock, Fallback), y la inyección en FastAPI.


El problema sin DI: acoplamiento profundo

# ❌ Sin DI: el negocio está acoplado a OpenAI

# src/domain/sentiment_service.py
from openai import OpenAI  # ← Importación de infrastructure en domain
import os

def analyze_sentiment(text: str) -> dict:
    client = OpenAI(api_key=os.getenv("OPENAI_API_KEY"))  # ← Creación interna
    
    response = client.chat.completions.create(
        model="gpt-4o-mini",       # ← Hardcoded
        messages=[...],
        temperature=0.7            # ← Magic number
    )
    
    return parse_result(response.choices[0].message.content)

# Problemas:
# 
# 1. Para testear:
#    → Necesitas patch("openai.OpenAI") o patch("openai.chat.completions.create")
#    → Si OpenAI cambia su API internamente, tu mock se rompe
#    → El test está atado a la implementación de OpenAI
#
# 2. Para cambiar de proveedor:
#    → Buscar todas las importaciones `from openai import OpenAI` en todo el codebase
#    → Reescribir la lógica de llamada (que es diferente entre OpenAI y Anthropic)
#    → Probar que nada se rompió (sin garantía si no tienes tests)
#
# 3. Para añadir un fallback:
#    → Envolver el call con try/except y duplicar la lógica para otro proveedor
#    → Difícil de testear

El Protocol: definir la interface

# src/infrastructure/llm_provider.py
from typing import Protocol, runtime_checkable

@runtime_checkable
class LLMProvider(Protocol):
    """
    Interface para LLM providers.
    
    @runtime_checkable permite usar isinstance() en runtime:
    >>> isinstance(my_provider, LLMProvider)
    True
    
    Cualquier clase que tenga el método complete() con esta firma
    satisface el Protocol SIN herencia explícita.
    
    Esto es "duck typing" con type safety:
    - OpenAIProvider no necesita `class OpenAIProvider(LLMProvider)`
    - Solo necesita implementar complete() con la firma correcta
    """
    
    def complete(self, messages: list[dict], **kwargs) -> str:
        """
        Envía mensajes al LLM y retorna el contenido de texto de la respuesta.
        
        Args:
            messages: Lista de mensajes en formato estándar
                     [{"role": "system", "content": "..."}, 
                      {"role": "user", "content": "..."}]
            **kwargs: Parámetros opcionales específicos del provider
        
        Returns:
            El contenido del mensaje de respuesta como string
        
        Raises:
            LLMProviderError: Si el provider no puede completar la request
                              (rate limit, timeout, invalid response, etc.)
        """
        ...


class LLMProviderError(Exception):
    """
    Error genérico del LLM provider.
    
    Encapsula errores específicos de cada proveedor (OpenAIError, etc.)
    en un error genérico que el domain puede manejar sin conocer
    qué proveedor se está usando.
    """
    def __init__(self, message: str, original_error: Exception = None):
        super().__init__(message)
        self.original_error = original_error

Implementación OpenAI

# src/infrastructure/openai_provider.py
import time
import structlog
from src.infrastructure.llm_provider import LLMProvider, LLMProviderError
from src.logging_config import calculate_cost

log = structlog.get_logger()

class OpenAIProvider:
    """
    Implementación de LLMProvider para la API de OpenAI.
    
    Encapsula TODO lo específico de OpenAI:
    - Formato de la request (messages, model, temperature, etc.)
    - Formato de la response (choices[0].message.content)
    - Uso y costo (usage.prompt_tokens, usage.completion_tokens)
    - Manejo de errores específicos de OpenAI
    - Logging de tokens, costo, latencia
    """
    
    def __init__(self, client, model: str, temperature: float, max_tokens: int,
                 seed: int = None):
        self._client = client
        self._model = model
        self._temperature = temperature
        self._max_tokens = max_tokens
        self._seed = seed
    
    def complete(self, messages: list[dict], **kwargs) -> str:
        """Realiza una llamada al API de OpenAI."""
        start = time.time()
        
        request_kwargs = {
            "model": self._model,
            "messages": messages,
            "temperature": self._temperature,
            "max_tokens": self._max_tokens,
        }
        if self._seed is not None:
            request_kwargs["seed"] = self._seed
        request_kwargs.update(kwargs)
        
        try:
            response = self._client.chat.completions.create(**request_kwargs)
            
            duration_ms = (time.time() - start) * 1000
            cost_usd = calculate_cost(
                self._model,
                response.usage.prompt_tokens,
                response.usage.completion_tokens
            )
            
            log.info(
                "llm_call_completed",
                model=self._model,
                input_tokens=response.usage.prompt_tokens,
                output_tokens=response.usage.completion_tokens,
                cost_usd=cost_usd,
                duration_ms=round(duration_ms, 1),
                finish_reason=response.choices[0].finish_reason
            )
            
            return response.choices[0].message.content
        
        except Exception as e:
            log.error(
                "llm_call_failed",
                model=self._model,
                error_type=type(e).__name__,
                duration_ms=round((time.time() - start) * 1000, 1)
            )
            raise LLMProviderError(
                f"OpenAI call failed: {type(e).__name__}: {str(e)}",
                original_error=e
            )
    
    @classmethod
    def from_settings(cls, settings) -> "OpenAIProvider":
        """Factory para crear desde Settings."""
        return cls(
            client=settings.create_openai_client(),
            model=settings.model,
            temperature=settings.temperature,
            max_tokens=settings.max_tokens,
            seed=settings.seed
        )

Implementación Mock para tests

# src/infrastructure/mock_provider.py
from src.infrastructure.llm_provider import LLMProvider
from typing import Union, Callable, Optional

class MockProvider:
    """
    Mock del LLMProvider para tests y desarrollo.
    
    Soporta tres modos:
    1. response fijo: siempre devuelve el mismo string
    2. callable: devuelve lo que retorne la función
    3. secuencia: devuelve diferentes respuestas en cada llamada
    
    Registra todas las llamadas para verificación en tests.
    """
    
    def __init__(
        self,
        response: Union[str, Callable[[list], str]] = None,
        responses: list[str] = None,
        raise_error: Exception = None
    ):
        """
        Args:
            response: String fijo o callable(messages) -> str
            responses: Lista de strings, usados en orden (circular)
            raise_error: Si se especifica, lanzar este error en complete()
        """
        self._response = response or '{"sentiment": "positive", "score": 0.8, "confidence": 0.9}'
        self._responses = responses
        self._raise_error = raise_error
        self._call_count = 0
        self.calls: list[list[dict]] = []  # Registro de todas las llamadas
    
    def complete(self, messages: list[dict], **kwargs) -> str:
        """Retorna la respuesta mock sin llamar a ninguna API."""
        self.calls.append(messages)
        self._call_count += 1
        
        if self._raise_error:
            raise self._raise_error
        
        if self._responses:
            # Retornar respuestas en secuencia (circular)
            idx = (self._call_count - 1) % len(self._responses)
            return self._responses[idx]
        
        if callable(self._response):
            return self._response(messages)
        
        return self._response
    
    @property
    def call_count(self) -> int:
        return self._call_count
    
    def was_called_with_system_message(self, content: str) -> bool:
        """Verifica si se llamó con un system message específico."""
        for call_messages in self.calls:
            for msg in call_messages:
                if msg.get("role") == "system" and content in msg.get("content", ""):
                    return True
        return False
    
    def get_last_user_message(self) -> Optional[str]:
        """Retorna el último mensaje del usuario."""
        if not self.calls:
            return None
        for msg in reversed(self.calls[-1]):
            if msg.get("role") == "user":
                return msg.get("content")
        return None

Implementación FallbackProvider

# src/infrastructure/fallback_provider.py
import structlog
from src.infrastructure.llm_provider import LLMProvider, LLMProviderError

log = structlog.get_logger()

class FallbackProvider:
    """
    Provider que intenta múltiples providers en orden.
    Si el primero falla, intenta el siguiente.
    
    Útil para:
    - Alta disponibilidad: OpenAI → Anthropic → Ollama local
    - Reducción de costos: gpt-4o-mini (falla) → gpt-3.5-turbo
    - Desarrollo: API real (falla con rate limit) → mock local
    """
    
    def __init__(self, providers: list[LLMProvider], name: str = "fallback"):
        if len(providers) < 2:
            raise ValueError("FallbackProvider requiere al menos 2 providers")
        self._providers = providers
        self._name = name
    
    def complete(self, messages: list[dict], **kwargs) -> str:
        """Intenta cada provider en orden hasta que uno funcione."""
        last_error = None
        
        for i, provider in enumerate(self._providers):
            try:
                result = provider.complete(messages, **kwargs)
                
                if i > 0:
                    log.warning(
                        "fallback_provider_used",
                        fallback_index=i,
                        provider_type=type(provider).__name__
                    )
                
                return result
            
            except LLMProviderError as e:
                last_error = e
                log.warning(
                    "provider_failed_trying_next",
                    provider_index=i,
                    provider_type=type(provider).__name__,
                    error=str(e)[:100],
                    has_next=i < len(self._providers) - 1
                )
        
        # Todos los providers fallaron
        raise LLMProviderError(
            f"All {len(self._providers)} providers failed. "
            f"Last error: {last_error}",
            original_error=last_error
        )

Dependency injection en FastAPI

# src/app/dependencies.py
from functools import lru_cache
from src.config import get_settings
from src.infrastructure.llm_provider import LLMProvider
from src.infrastructure.openai_provider import OpenAIProvider
from src.infrastructure.mock_provider import MockProvider

def get_llm_provider() -> LLMProvider:
    """
    FastAPI dependency que crea y retorna el LLM provider correcto
    basado en la configuración actual.
    
    Este es el único lugar donde se decide qué implementación usar.
    El resto del código (domain, endpoints) solo conoce LLMProvider.
    """
    settings = get_settings()
    
    if settings.use_mock_llm:
        return MockProvider(response=settings.mock_response)
    
    return OpenAIProvider.from_settings(settings)

# Para singleton (el mismo provider para toda la app):
@lru_cache(maxsize=1)
def get_cached_provider() -> LLMProvider:
    """Provider singleton — se crea una vez y se reutiliza."""
    return get_llm_provider()


# src/app/routers/sentiment.py
from fastapi import APIRouter, Depends
from src.domain.sentiment_service import analyze_sentiment
from src.infrastructure.llm_provider import LLMProvider
from src.app.dependencies import get_llm_provider

router = APIRouter()

@router.post("/analyze")
async def analyze_endpoint(
    body: AnalyzeRequest,
    provider: LLMProvider = Depends(get_llm_provider)  # Inyectado por FastAPI
):
    """
    El endpoint no sabe si está usando OpenAI, Anthropic, o Mock.
    Solo sabe que tiene un LLMProvider.
    """
    result = analyze_sentiment(text=body.text, provider=provider)
    return AnalyzeResponse(**result, request_id=get_request_id())

Tests simples gracias a DI

# tests/unit/test_sentiment_service.py
import pytest
from src.domain.sentiment_service import analyze_sentiment, LowConfidenceError
from src.infrastructure.mock_provider import MockProvider

class TestAnalyzeSentiment:
    """
    Tests del domain usando MockProvider.
    
    No hay:
    - patch() de ningún módulo
    - API keys de OpenAI
    - Conexión a internet
    - Sleep o rate limits
    
    Cada test es determinístico y rápido (<10ms).
    """
    
    def test_positive_sentiment(self):
        mock = MockProvider('{"sentiment": "positive", "score": 0.8, "confidence": 0.9}')
        result = analyze_sentiment("This is great!", mock)
        
        assert result["sentiment"] == "positive"
        assert result["score"] == 0.8
        assert result["confidence"] == 0.9
    
    def test_negative_sentiment(self):
        mock = MockProvider('{"sentiment": "negative", "score": -0.7, "confidence": 0.85}')
        result = analyze_sentiment("This is terrible", mock)
        
        assert result["sentiment"] == "negative"
    
    def test_low_confidence_raises_error(self):
        mock = MockProvider('{"sentiment": "mixed", "score": 0.1, "confidence": 0.1}')
        
        with pytest.raises(LowConfidenceError):
            analyze_sentiment("ambiguous text", mock)
    
    def test_invalid_json_from_llm_returns_unknown(self):
        mock = MockProvider("I'm sorry, I cannot analyze this text.")
        result = analyze_sentiment("test", mock)
        
        assert result["sentiment"] == "unknown"
        assert result["score"] == 0.0
    
    def test_provider_called_once(self):
        """Verifica que se llama al provider exactamente una vez."""
        mock = MockProvider('{"sentiment": "positive", "score": 0.9, "confidence": 0.95}')
        
        analyze_sentiment("test text", mock)
        
        assert mock.call_count == 1
    
    def test_prompt_contains_input_text(self):
        """Verifica que el texto del usuario está en el prompt enviado al LLM."""
        mock = MockProvider('{"sentiment": "positive", "score": 0.9, "confidence": 0.95}')
        test_text = "unique_test_string_12345"
        
        analyze_sentiment(test_text, mock)
        
        last_message = mock.get_last_user_message()
        assert test_text in last_message, \
            f"Expected '{test_text}' in the prompt, got: {last_message}"
    
    def test_fallback_provider_tries_secondary_on_failure(self):
        """FallbackProvider usa el segundo si el primero falla."""
        from src.infrastructure.llm_provider import LLMProviderError
        from src.infrastructure.fallback_provider import FallbackProvider
        
        failing_provider = MockProvider(
            raise_error=LLMProviderError("Rate limit exceeded")
        )
        working_provider = MockProvider(
            '{"sentiment": "positive", "score": 0.9, "confidence": 0.95}'
        )
        
        fallback = FallbackProvider([failing_provider, working_provider])
        result = analyze_sentiment("test", fallback)
        
        assert result["sentiment"] == "positive"
        assert failing_provider.call_count == 1  # Intentó el primero
        assert working_provider.call_count == 1  # Usó el segundo

Comparación completa: sin vs con DI

Sin DI:

TESTEAR analyze_sentiment():
  → patch("openai.OpenAI") para mockear el cliente
  → patch("openai.chat.completions.create") para la respuesta
  → Configurar el mock con la estructura exacta de OpenAI API
  → Si OpenAI cambia su API, el mock se rompe aunque tu código esté bien
  → Tiempo de setup: 15-20 líneas de código por test

Con DI:

TESTEAR analyze_sentiment():
  → mock = MockProvider(response='{"sentiment": "positive", ...}')
  → result = analyze_sentiment("text", mock)
  → Tiempo de setup: 1 línea de código por test

─────────────────────────────────────────────

Sin DI:

CAMBIAR OpenAI → Anthropic:
  → Buscar todos los `from openai import` en el proyecto
  → Buscar todos los `client.chat.completions.create()`
  → Reescribir con la API de Anthropic (formato diferente)
  → Testear todo manualmente
  → Tiempo: 1-2 días

Con DI:

CAMBIAR OpenAI → Anthropic:
  → Crear AnthropicProvider(LLMProvider) con su complete()
  → En dependencies.py: retornar AnthropicProvider() en vez de OpenAIProvider()
  → Todos los tests del domain siguen pasando sin cambios
  → Tiempo: 2-4 horas

Ejercicios

Ejercicio 1: Implementar AnthropicProvider

Escribe el esqueleto de AnthropicProvider que implementa LLMProvider:

Ver solución
# src/infrastructure/anthropic_provider.py
from src.infrastructure.llm_provider import LLMProvider, LLMProviderError

class AnthropicProvider:
    """Implementación de LLMProvider para Anthropic Claude."""
    
    def __init__(self, client, model: str = "claude-3-haiku-20240307",
                 max_tokens: int = 500):
        self._client = client
        self._model = model
        self._max_tokens = max_tokens
    
    def complete(self, messages: list[dict], **kwargs) -> str:
        # Anthropic tiene formato diferente — system message separado
        system = ""
        user_messages = []
        for msg in messages:
            if msg["role"] == "system":
                system = msg["content"]
            else:
                user_messages.append(msg)
        
        try:
            response = self._client.messages.create(
                model=self._model,
                max_tokens=self._max_tokens,
                system=system,
                messages=user_messages
            )
            return response.content[0].text
        except Exception as e:
            raise LLMProviderError(f"Anthropic call failed: {e}", original_error=e)

Ejercicio 2: Test de FallbackProvider

Escribe un test que verifica que si el primer provider lanza LLMProviderError 3 veces seguidas, el FallbackProvider usa el segundo:

Ver solución
def test_fallback_after_multiple_failures():
    from src.infrastructure.llm_provider import LLMProviderError
    from src.infrastructure.fallback_provider import FallbackProvider
    
    primary = MockProvider(raise_error=LLMProviderError("Primary unavailable"))
    secondary = MockProvider('{"sentiment": "neutral", "score": 0.0, "confidence": 0.8}')
    
    fallback = FallbackProvider([primary, secondary])
    
    # Tres llamadas separadas
    for _ in range(3):
        result = fallback.complete([{"role": "user", "content": "test"}])
        data = json.loads(result)
        assert data["sentiment"] == "neutral"
    
    assert primary.call_count == 3    # Se intentó 3 veces
    assert secondary.call_count == 3  # Se usó como fallback 3 veces

Resumen

  • Protocol define la interface que cualquier LLM provider debe implementar — sin herencia explícita
  • OpenAIProvider encapsula toda la lógica específica de OpenAI (formato, logging, errores)
  • MockProvider configurable para tests: respuesta fija, callable, secuencia, o error
  • FallbackProvider implementa alta disponibilidad: si A falla, intenta B
  • Depends(get_llm_provider) en FastAPI: el endpoint recibe un LLMProvider sin saber cuál
  • Tests sin patch: con DI, los unit tests del domain son triviales y no se rompen con cambios de API

Recursos adicionales

  1. typing.Protocol Python Docs — Documentación oficial
  2. FastAPI Dependency Injection — DI en FastAPI
  3. Dependency Injection (Martin Fowler) — El artículo original
  4. Hexagonal Architecture (Alistair Cockburn) — Ports and Adapters
  5. Anthropic Python SDK — Para implementar AnthropicProvider