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
- typing.Protocol Python Docs — Documentación oficial
- FastAPI Dependency Injection — DI en FastAPI
- Dependency Injection (Martin Fowler) — El artículo original
- Hexagonal Architecture (Alistair Cockburn) — Ports and Adapters
- Anthropic Python SDK — Para implementar AnthropicProvider