Files
forja/tests/unit/test_llm_factory.py
2026-06-10 18:21:33 +02:00

68 lines
2.2 KiB
Python

"""Tests del factory LLM."""
from typing import Any
import pytest
from forja_core.config import Settings
from forja_core.llm.base import CompletionResult, Message
from forja_core.llm.factory import build_llm_provider
from forja_core.llm.fallback import FallbackLLMProvider
from forja_core.llm.mock import MockProvider
def test_factory_devuelve_mock_por_defecto() -> None:
s = Settings(llm_provider="mock", _env_file=None) # type: ignore[call-arg]
p = build_llm_provider(s)
assert isinstance(p, MockProvider)
assert p.name == "mock"
def test_factory_azure_requiere_credenciales() -> None:
s = Settings(llm_provider="azure", _env_file=None) # type: ignore[call-arg]
with pytest.raises(ValueError):
build_llm_provider(s)
def test_factory_envuelve_con_fallback_cuando_difiere() -> None:
s = Settings(
llm_provider="openai",
llm_fallback_provider="mock",
openai_api_key="sk-test",
_env_file=None,
) # type: ignore[call-arg]
p = build_llm_provider(s)
assert isinstance(p, FallbackLLMProvider)
assert p.name == "openai+fallback:mock"
def test_factory_ignora_fallback_igual_al_primario() -> None:
s = Settings(llm_provider="mock", llm_fallback_provider="mock", _env_file=None) # type: ignore[call-arg]
p = build_llm_provider(s)
assert isinstance(p, MockProvider)
class _FailingProvider:
name = "failing"
async def complete(
self,
messages: list[Message],
schema: dict[str, Any] | None = None,
temperature: float = 0.2,
max_tokens: int = 2000,
) -> CompletionResult:
raise RuntimeError("primary down")
async def test_fallback_delega_cuando_el_primario_falla() -> None:
provider = FallbackLLMProvider(_FailingProvider(), MockProvider())
result = await provider.complete([Message(role="user", content="mos degradation")])
assert result.model.startswith("mock")
async def test_fallback_no_interviene_si_el_primario_responde() -> None:
provider = FallbackLLMProvider(MockProvider(), _FailingProvider())
result = await provider.complete([Message(role="user", content="mos degradation")])
assert result.model.startswith("mock")