202 lines
6.7 KiB
Python
202 lines
6.7 KiB
Python
"""Tests for LLMProviderChain."""
|
|
|
|
from unittest.mock import AsyncMock
|
|
|
|
import httpx
|
|
import pytest
|
|
|
|
from web.config import WebSettings
|
|
from web.exceptions import ProviderError
|
|
from web.llm.provider import LLMProviderChain, _convert_to_anthropic_messages
|
|
|
|
|
|
def _make_httpx_response(status_code: int, json: dict) -> httpx.Response:
|
|
"""Create an httpx.Response with a request set (needed for raise_for_status)."""
|
|
response = httpx.Response(status_code, json=json)
|
|
response._request = httpx.Request("POST", "http://test")
|
|
return response
|
|
|
|
|
|
@pytest.fixture
|
|
def settings() -> WebSettings:
|
|
"""Create test settings."""
|
|
return WebSettings(
|
|
searxng_base_url="http://localhost:55100",
|
|
llm_base_url="http://localhost:14011",
|
|
external_url="http://localhost:51100",
|
|
openai_api_key="test-openai-key",
|
|
anthropic_api_key="test-anthropic-key",
|
|
)
|
|
|
|
|
|
@pytest.fixture
|
|
def settings_no_keys() -> WebSettings:
|
|
"""Create test settings without API keys."""
|
|
return WebSettings(
|
|
searxng_base_url="http://localhost:55100",
|
|
llm_base_url="http://localhost:14011",
|
|
external_url="http://localhost:51100",
|
|
)
|
|
|
|
|
|
class TestLLMProviderChainFallback:
|
|
"""Tests for provider fallback behavior."""
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_auto_tries_local_first(self, settings: WebSettings) -> None:
|
|
"""Test that auto provider tries local first."""
|
|
mock_client = AsyncMock()
|
|
mock_client.post = AsyncMock(
|
|
return_value=_make_httpx_response(
|
|
200,
|
|
{
|
|
"choices": [{"message": {"content": "Hello"}}],
|
|
"usage": {"total_tokens": 10},
|
|
},
|
|
)
|
|
)
|
|
|
|
async def get_client():
|
|
return mock_client
|
|
|
|
chain = LLMProviderChain(settings, get_client)
|
|
text, tokens, provider, _model = await chain.call_chat(
|
|
messages=[{"role": "user", "content": "Hi"}],
|
|
provider="auto",
|
|
max_tokens=100,
|
|
)
|
|
|
|
assert provider == "local"
|
|
assert text == "Hello"
|
|
assert tokens == 10
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_auto_falls_back_to_openai(self, settings: WebSettings) -> None:
|
|
"""Test fallback to OpenAI when local fails."""
|
|
mock_client = AsyncMock()
|
|
|
|
# First call (local) fails, second (openai) succeeds
|
|
mock_client.post = AsyncMock(
|
|
side_effect=[
|
|
httpx.ConnectError("refused"),
|
|
_make_httpx_response(
|
|
200,
|
|
{
|
|
"choices": [{"message": {"content": "From OpenAI"}}],
|
|
"usage": {"total_tokens": 20},
|
|
},
|
|
),
|
|
]
|
|
)
|
|
|
|
async def get_client():
|
|
return mock_client
|
|
|
|
chain = LLMProviderChain(settings, get_client)
|
|
text, _tokens, provider, _model = await chain.call_chat(
|
|
messages=[{"role": "user", "content": "Hi"}],
|
|
provider="auto",
|
|
max_tokens=100,
|
|
)
|
|
|
|
assert provider == "openai"
|
|
assert text == "From OpenAI"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_skips_providers_without_keys(
|
|
self, settings_no_keys: WebSettings
|
|
) -> None:
|
|
"""Test that providers without API keys are skipped."""
|
|
mock_client = AsyncMock()
|
|
mock_client.post = AsyncMock(
|
|
return_value=_make_httpx_response(
|
|
200,
|
|
{
|
|
"choices": [{"message": {"content": "Local"}}],
|
|
"usage": {"total_tokens": 5},
|
|
},
|
|
)
|
|
)
|
|
|
|
async def get_client():
|
|
return mock_client
|
|
|
|
chain = LLMProviderChain(settings_no_keys, get_client)
|
|
_text, _tokens, provider, _model = await chain.call_chat(
|
|
messages=[{"role": "user", "content": "Hi"}],
|
|
provider="auto",
|
|
max_tokens=100,
|
|
)
|
|
|
|
assert provider == "local"
|
|
# Only one call made (openai and anthropic skipped due to missing keys)
|
|
assert mock_client.post.call_count == 1
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_raises_when_all_fail(self, settings_no_keys: WebSettings) -> None:
|
|
"""Test ProviderError when all providers fail."""
|
|
mock_client = AsyncMock()
|
|
mock_client.post = AsyncMock(side_effect=httpx.ConnectError("refused"))
|
|
|
|
async def get_client():
|
|
return mock_client
|
|
|
|
chain = LLMProviderChain(settings_no_keys, get_client, domain="test")
|
|
with pytest.raises(ProviderError, match="test"):
|
|
await chain.call_chat(
|
|
messages=[{"role": "user", "content": "Hi"}],
|
|
provider="auto",
|
|
max_tokens=100,
|
|
)
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_raises_no_provider_available(
|
|
self, settings_no_keys: WebSettings
|
|
) -> None:
|
|
"""Test ProviderError when no providers are available (all skipped)."""
|
|
mock_client = AsyncMock()
|
|
|
|
async def get_client():
|
|
return mock_client
|
|
|
|
chain = LLMProviderChain(settings_no_keys, get_client, domain="vision")
|
|
with pytest.raises(ProviderError, match="No vision provider available"):
|
|
await chain.call_chat(
|
|
messages=[{"role": "user", "content": "Hi"}],
|
|
provider="openai", # No key set
|
|
max_tokens=100,
|
|
)
|
|
|
|
|
|
class TestAnthropicMessageConversion:
|
|
"""Tests for OpenAI -> Anthropic message format conversion."""
|
|
|
|
def test_text_only_message(self) -> None:
|
|
"""Test conversion of text-only messages."""
|
|
messages = [{"role": "user", "content": "Hello"}]
|
|
result = _convert_to_anthropic_messages(messages)
|
|
assert result == [{"role": "user", "content": "Hello"}]
|
|
|
|
def test_multimodal_with_base64_image(self) -> None:
|
|
"""Test conversion of messages with base64 images."""
|
|
messages = [
|
|
{
|
|
"role": "user",
|
|
"content": [
|
|
{
|
|
"type": "image_url",
|
|
"image_url": {"url": "data:image/jpeg;base64,AAAA"},
|
|
},
|
|
{"type": "text", "text": "Describe this."},
|
|
],
|
|
}
|
|
]
|
|
result = _convert_to_anthropic_messages(messages)
|
|
assert len(result) == 1
|
|
content = result[0]["content"]
|
|
assert len(content) == 2
|
|
assert content[0]["type"] == "image"
|
|
assert content[0]["source"]["type"] == "base64"
|
|
assert content[0]["source"]["media_type"] == "image/jpeg"
|
|
assert content[0]["source"]["data"] == "AAAA"
|
|
assert content[1]["type"] == "text"
|