FastAPI + Celery + Next.js + Postgres/Redis app with company monitoring, source collection, LLM-based change analysis, enrichment, and account security (Turnstile, escalating lockout, email verification).
180 lines
5.9 KiB
Python
180 lines
5.9 KiB
Python
"""Anthropic/Ollama providers, with the SDK/HTTP layer mocked - these never
|
|
run against a real paid API in the test suite. Verifies the structured-
|
|
output + repair-loop wiring actually works, not just that it imports."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
from types import SimpleNamespace
|
|
from unittest.mock import AsyncMock, patch
|
|
|
|
import httpx
|
|
import pytest
|
|
import respx
|
|
from pydantic import BaseModel
|
|
|
|
from app.analysis.llm.base import LLMResponseError
|
|
from app.core.config import Settings
|
|
|
|
|
|
class _Toy(BaseModel):
|
|
answer: str
|
|
|
|
|
|
def _settings(**overrides) -> Settings:
|
|
defaults = {
|
|
"llm_provider": "anthropic",
|
|
"anthropic_api_key": "sk-test",
|
|
"anthropic_model": "claude-test",
|
|
"llm_max_retries": 1,
|
|
"llm_max_tokens_per_request": 100,
|
|
}
|
|
defaults.update(overrides)
|
|
return Settings(**defaults)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_anthropic_provider_parses_tool_use_response():
|
|
from app.analysis.llm.anthropic_provider import AnthropicLLMProvider
|
|
|
|
provider = AnthropicLLMProvider(_settings())
|
|
fake_response = SimpleNamespace(
|
|
content=[SimpleNamespace(type="tool_use", input={"answer": "42"})]
|
|
)
|
|
with patch.object(provider._client.messages, "create", AsyncMock(return_value=fake_response)):
|
|
result = await provider.generate_structured("system", "user", _Toy)
|
|
|
|
assert result.answer == "42"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_anthropic_provider_retries_then_raises_on_missing_tool_use():
|
|
from app.analysis.llm.anthropic_provider import AnthropicLLMProvider
|
|
|
|
provider = AnthropicLLMProvider(_settings(llm_max_retries=1))
|
|
fake_response = SimpleNamespace(content=[SimpleNamespace(type="text", text="oops")])
|
|
with patch.object(
|
|
provider._client.messages, "create", AsyncMock(return_value=fake_response)
|
|
) as mock_create:
|
|
with pytest.raises(LLMResponseError):
|
|
await provider.generate_structured("system", "user", _Toy)
|
|
|
|
assert mock_create.call_count == 2 # initial attempt + 1 retry
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_anthropic_provider_generate_text_joins_text_blocks():
|
|
from app.analysis.llm.anthropic_provider import AnthropicLLMProvider
|
|
|
|
provider = AnthropicLLMProvider(_settings())
|
|
fake_response = SimpleNamespace(
|
|
content=[
|
|
SimpleNamespace(type="text", text="Hello"),
|
|
SimpleNamespace(type="text", text="world"),
|
|
]
|
|
)
|
|
with patch.object(provider._client.messages, "create", AsyncMock(return_value=fake_response)):
|
|
text = await provider.generate_text("system", "user")
|
|
|
|
assert text == "Hello\nworld"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_ollama_provider_parses_json_response():
|
|
from app.analysis.llm.ollama_provider import OllamaLLMProvider
|
|
|
|
settings = Settings(
|
|
llm_provider="ollama",
|
|
ollama_base_url="http://ollama.local:11434",
|
|
ollama_model="llama-test",
|
|
llm_max_retries=1,
|
|
)
|
|
provider = OllamaLLMProvider(settings)
|
|
|
|
with respx.mock:
|
|
respx.post("http://ollama.local:11434/api/chat").mock(
|
|
return_value=httpx.Response(
|
|
200, json={"message": {"content": json.dumps({"answer": "42"})}}
|
|
)
|
|
)
|
|
result = await provider.generate_structured("system", "user", _Toy)
|
|
|
|
assert result.answer == "42"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_ollama_provider_retries_then_raises_on_invalid_json():
|
|
from app.analysis.llm.ollama_provider import OllamaLLMProvider
|
|
|
|
settings = Settings(
|
|
llm_provider="ollama",
|
|
ollama_base_url="http://ollama.local:11434",
|
|
ollama_model="llama-test",
|
|
llm_max_retries=1,
|
|
)
|
|
provider = OllamaLLMProvider(settings)
|
|
|
|
with respx.mock:
|
|
route = respx.post("http://ollama.local:11434/api/chat").mock(
|
|
return_value=httpx.Response(200, json={"message": {"content": "not json"}})
|
|
)
|
|
with pytest.raises(LLMResponseError):
|
|
await provider.generate_structured("system", "user", _Toy)
|
|
|
|
assert route.call_count == 2
|
|
|
|
|
|
def _gemini_settings(**overrides) -> Settings:
|
|
defaults = {
|
|
"llm_provider": "gemini",
|
|
"gemini_api_key": "test-key",
|
|
"gemini_model": "gemini-test",
|
|
"llm_max_retries": 1,
|
|
"llm_max_tokens_per_request": 100,
|
|
}
|
|
defaults.update(overrides)
|
|
return Settings(**defaults)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_gemini_provider_parses_structured_response():
|
|
from app.analysis.llm.gemini_provider import GeminiLLMProvider
|
|
|
|
provider = GeminiLLMProvider(_gemini_settings())
|
|
fake_response = SimpleNamespace(parsed=_Toy(answer="42"))
|
|
with patch.object(
|
|
provider._client.aio.models, "generate_content", AsyncMock(return_value=fake_response)
|
|
):
|
|
result = await provider.generate_structured("system", "user", _Toy)
|
|
|
|
assert result.answer == "42"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_gemini_provider_retries_then_raises_when_unparsed():
|
|
from app.analysis.llm.gemini_provider import GeminiLLMProvider
|
|
|
|
provider = GeminiLLMProvider(_gemini_settings(llm_max_retries=1))
|
|
fake_response = SimpleNamespace(parsed=None)
|
|
with patch.object(
|
|
provider._client.aio.models, "generate_content", AsyncMock(return_value=fake_response)
|
|
) as mock_generate:
|
|
with pytest.raises(LLMResponseError):
|
|
await provider.generate_structured("system", "user", _Toy)
|
|
|
|
assert mock_generate.call_count == 2 # initial attempt + 1 retry
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_gemini_provider_generate_text_returns_response_text():
|
|
from app.analysis.llm.gemini_provider import GeminiLLMProvider
|
|
|
|
provider = GeminiLLMProvider(_gemini_settings())
|
|
fake_response = SimpleNamespace(text="Hello world")
|
|
with patch.object(
|
|
provider._client.aio.models, "generate_content", AsyncMock(return_value=fake_response)
|
|
):
|
|
text = await provider.generate_text("system", "user")
|
|
|
|
assert text == "Hello world"
|