feat: add AI provider abstraction layer with OpenAI, Azure, Anthropic, Gemini, Ollama, OpenRouter, LiteLLM support
Co-authored-by: christianlouis <361235+christianlouis@users.noreply.github.com>
This commit is contained in:
@@ -0,0 +1,591 @@
|
||||
"""Unit tests for the AI provider abstraction layer (app/utils/ai_provider.py)."""
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from app.utils.ai_provider import (
|
||||
AIProvider,
|
||||
AnthropicProvider,
|
||||
AzureOpenAIProvider,
|
||||
GeminiProvider,
|
||||
LiteLLMProvider,
|
||||
OllamaProvider,
|
||||
OpenAIProvider,
|
||||
OpenRouterProvider,
|
||||
get_ai_provider,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _make_openai_response(content: str) -> MagicMock:
|
||||
"""Build a minimal mock that looks like an openai ChatCompletion response."""
|
||||
msg = SimpleNamespace(content=content)
|
||||
choice = SimpleNamespace(message=msg)
|
||||
resp = MagicMock()
|
||||
resp.choices = [choice]
|
||||
return resp
|
||||
|
||||
|
||||
def _make_litellm_response(content: str) -> MagicMock:
|
||||
"""Build a minimal mock that looks like a litellm completion response."""
|
||||
msg = SimpleNamespace(content=content)
|
||||
choice = SimpleNamespace(message=msg)
|
||||
resp = MagicMock()
|
||||
resp.choices = [choice]
|
||||
return resp
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Abstract base class
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestAIProviderABC:
|
||||
"""Verify that AIProvider cannot be instantiated directly."""
|
||||
|
||||
def test_cannot_instantiate_directly(self):
|
||||
"""AIProvider is abstract and must be subclassed."""
|
||||
with pytest.raises(TypeError):
|
||||
AIProvider() # type: ignore[abstract]
|
||||
|
||||
def test_subclass_must_implement_chat_completion(self):
|
||||
"""Concrete subclass without chat_completion raises TypeError."""
|
||||
|
||||
class Incomplete(AIProvider):
|
||||
pass
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
Incomplete() # type: ignore[abstract]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# OpenAIProvider
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestOpenAIProvider:
|
||||
"""Tests for OpenAIProvider."""
|
||||
|
||||
@patch("openai.OpenAI")
|
||||
def test_initialises_with_default_base_url(self, mock_openai_cls):
|
||||
"""Provider uses the OpenAI default endpoint when no base_url is given."""
|
||||
OpenAIProvider(api_key="sk-test")
|
||||
mock_openai_cls.assert_called_once_with(
|
||||
api_key="sk-test",
|
||||
base_url="https://api.openai.com/v1",
|
||||
)
|
||||
|
||||
@patch("openai.OpenAI")
|
||||
def test_initialises_with_custom_base_url(self, mock_openai_cls):
|
||||
"""Provider passes through a custom base_url."""
|
||||
OpenAIProvider(api_key="sk-test", base_url="http://localhost:8000/v1")
|
||||
mock_openai_cls.assert_called_once_with(
|
||||
api_key="sk-test",
|
||||
base_url="http://localhost:8000/v1",
|
||||
)
|
||||
|
||||
@patch("openai.OpenAI")
|
||||
def test_chat_completion_returns_content(self, mock_openai_cls):
|
||||
"""chat_completion extracts and returns the message content string."""
|
||||
mock_client = MagicMock()
|
||||
mock_client.chat.completions.create.return_value = _make_openai_response("Hello world")
|
||||
mock_openai_cls.return_value = mock_client
|
||||
|
||||
provider = OpenAIProvider(api_key="sk-test")
|
||||
result = provider.chat_completion(
|
||||
messages=[{"role": "user", "content": "Hi"}],
|
||||
model="gpt-4o-mini",
|
||||
temperature=0,
|
||||
)
|
||||
|
||||
assert result == "Hello world"
|
||||
mock_client.chat.completions.create.assert_called_once_with(
|
||||
model="gpt-4o-mini",
|
||||
messages=[{"role": "user", "content": "Hi"}],
|
||||
temperature=0,
|
||||
)
|
||||
|
||||
@patch("openai.OpenAI")
|
||||
def test_chat_completion_passes_kwargs(self, mock_openai_cls):
|
||||
"""Extra kwargs are forwarded to the underlying client."""
|
||||
mock_client = MagicMock()
|
||||
mock_client.chat.completions.create.return_value = _make_openai_response("ok")
|
||||
mock_openai_cls.return_value = mock_client
|
||||
|
||||
provider = OpenAIProvider(api_key="sk-test")
|
||||
provider.chat_completion(
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
model="gpt-4o",
|
||||
temperature=0.7,
|
||||
max_tokens=100,
|
||||
)
|
||||
|
||||
call_kwargs = mock_client.chat.completions.create.call_args[1]
|
||||
assert call_kwargs["max_tokens"] == 100
|
||||
assert call_kwargs["temperature"] == 0.7
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# AzureOpenAIProvider
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestAzureOpenAIProvider:
|
||||
"""Tests for AzureOpenAIProvider."""
|
||||
|
||||
@patch("openai.AzureOpenAI")
|
||||
def test_initialises_correctly(self, mock_azure_cls):
|
||||
"""Azure provider initialises with correct parameters."""
|
||||
AzureOpenAIProvider(
|
||||
api_key="azure-key",
|
||||
azure_endpoint="https://my-resource.openai.azure.com",
|
||||
api_version="2024-02-01",
|
||||
)
|
||||
mock_azure_cls.assert_called_once_with(
|
||||
api_key="azure-key",
|
||||
azure_endpoint="https://my-resource.openai.azure.com",
|
||||
api_version="2024-02-01",
|
||||
)
|
||||
|
||||
@patch("openai.AzureOpenAI")
|
||||
def test_chat_completion_returns_content(self, mock_azure_cls):
|
||||
"""chat_completion returns the message content string."""
|
||||
mock_client = MagicMock()
|
||||
mock_client.chat.completions.create.return_value = _make_openai_response("Azure response")
|
||||
mock_azure_cls.return_value = mock_client
|
||||
|
||||
provider = AzureOpenAIProvider(
|
||||
api_key="key",
|
||||
azure_endpoint="https://endpoint.openai.azure.com",
|
||||
)
|
||||
result = provider.chat_completion(
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
model="gpt-4",
|
||||
)
|
||||
|
||||
assert result == "Azure response"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# AnthropicProvider
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestAnthropicProvider:
|
||||
"""Tests for AnthropicProvider."""
|
||||
|
||||
@patch("litellm.completion")
|
||||
def test_chat_completion_adds_anthropic_prefix(self, mock_completion):
|
||||
"""Provider prepends 'anthropic/' to bare model names."""
|
||||
mock_completion.return_value = _make_litellm_response("Claude says hi")
|
||||
|
||||
provider = AnthropicProvider(api_key="ant-key")
|
||||
result = provider.chat_completion(
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
model="claude-3-5-sonnet-20241022",
|
||||
)
|
||||
|
||||
assert result == "Claude says hi"
|
||||
call_kwargs = mock_completion.call_args[1]
|
||||
assert call_kwargs["model"] == "anthropic/claude-3-5-sonnet-20241022"
|
||||
|
||||
@patch("litellm.completion")
|
||||
def test_chat_completion_keeps_existing_prefix(self, mock_completion):
|
||||
"""Provider does not double-add 'anthropic/' if already present."""
|
||||
mock_completion.return_value = _make_litellm_response("ok")
|
||||
|
||||
provider = AnthropicProvider(api_key="ant-key")
|
||||
provider.chat_completion(
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
model="anthropic/claude-3-opus-20240229",
|
||||
)
|
||||
|
||||
call_kwargs = mock_completion.call_args[1]
|
||||
assert call_kwargs["model"] == "anthropic/claude-3-opus-20240229"
|
||||
|
||||
@patch("litellm.completion")
|
||||
def test_chat_completion_passes_api_key(self, mock_completion):
|
||||
"""Provider passes the API key to litellm."""
|
||||
mock_completion.return_value = _make_litellm_response("ok")
|
||||
|
||||
provider = AnthropicProvider(api_key="my-ant-key")
|
||||
provider.chat_completion(
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
model="claude-3-5-sonnet-20241022",
|
||||
)
|
||||
|
||||
call_kwargs = mock_completion.call_args[1]
|
||||
assert call_kwargs["api_key"] == "my-ant-key"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GeminiProvider
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestGeminiProvider:
|
||||
"""Tests for GeminiProvider."""
|
||||
|
||||
@patch("litellm.completion")
|
||||
def test_chat_completion_adds_gemini_prefix(self, mock_completion):
|
||||
"""Provider prepends 'gemini/' to bare model names."""
|
||||
mock_completion.return_value = _make_litellm_response("Gemini says hi")
|
||||
|
||||
provider = GeminiProvider(api_key="gemini-key")
|
||||
result = provider.chat_completion(
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
model="gemini-1.5-pro",
|
||||
)
|
||||
|
||||
assert result == "Gemini says hi"
|
||||
call_kwargs = mock_completion.call_args[1]
|
||||
assert call_kwargs["model"] == "gemini/gemini-1.5-pro"
|
||||
|
||||
@patch("litellm.completion")
|
||||
def test_chat_completion_keeps_existing_prefix(self, mock_completion):
|
||||
"""Provider does not double-add 'gemini/' if already present."""
|
||||
mock_completion.return_value = _make_litellm_response("ok")
|
||||
|
||||
provider = GeminiProvider(api_key="gemini-key")
|
||||
provider.chat_completion(
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
model="gemini/gemini-pro",
|
||||
)
|
||||
|
||||
call_kwargs = mock_completion.call_args[1]
|
||||
assert call_kwargs["model"] == "gemini/gemini-pro"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# OllamaProvider
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestOllamaProvider:
|
||||
"""Tests for OllamaProvider."""
|
||||
|
||||
@patch("openai.OpenAI")
|
||||
def test_initialises_with_default_base_url(self, mock_openai_cls):
|
||||
"""Provider appends /v1 to the default Ollama base URL."""
|
||||
OllamaProvider()
|
||||
mock_openai_cls.assert_called_once_with(
|
||||
api_key="ollama",
|
||||
base_url="http://localhost:11434/v1",
|
||||
)
|
||||
|
||||
@patch("openai.OpenAI")
|
||||
def test_initialises_with_custom_base_url(self, mock_openai_cls):
|
||||
"""Provider appends /v1 to a custom Ollama base URL."""
|
||||
OllamaProvider(base_url="http://my-ollama:11434")
|
||||
mock_openai_cls.assert_called_once_with(
|
||||
api_key="ollama",
|
||||
base_url="http://my-ollama:11434/v1",
|
||||
)
|
||||
|
||||
@patch("openai.OpenAI")
|
||||
def test_strips_trailing_slash_before_appending_v1(self, mock_openai_cls):
|
||||
"""Provider normalises trailing slashes before appending /v1."""
|
||||
OllamaProvider(base_url="http://ollama:11434/")
|
||||
call_kwargs = mock_openai_cls.call_args[1]
|
||||
assert call_kwargs["base_url"] == "http://ollama:11434/v1"
|
||||
|
||||
@patch("openai.OpenAI")
|
||||
def test_chat_completion_returns_content(self, mock_openai_cls):
|
||||
"""chat_completion returns the message content string."""
|
||||
mock_client = MagicMock()
|
||||
mock_client.chat.completions.create.return_value = _make_openai_response("Llama response")
|
||||
mock_openai_cls.return_value = mock_client
|
||||
|
||||
provider = OllamaProvider()
|
||||
result = provider.chat_completion(
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
model="llama3.2",
|
||||
)
|
||||
|
||||
assert result == "Llama response"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# OpenRouterProvider
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestOpenRouterProvider:
|
||||
"""Tests for OpenRouterProvider."""
|
||||
|
||||
@patch("openai.OpenAI")
|
||||
def test_initialises_with_default_base_url(self, mock_openai_cls):
|
||||
"""Provider uses the default OpenRouter endpoint."""
|
||||
OpenRouterProvider(api_key="or-key")
|
||||
mock_openai_cls.assert_called_once_with(
|
||||
api_key="or-key",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
)
|
||||
|
||||
@patch("openai.OpenAI")
|
||||
def test_chat_completion_returns_content(self, mock_openai_cls):
|
||||
"""chat_completion returns the message content string."""
|
||||
mock_client = MagicMock()
|
||||
mock_client.chat.completions.create.return_value = _make_openai_response("Router response")
|
||||
mock_openai_cls.return_value = mock_client
|
||||
|
||||
provider = OpenRouterProvider(api_key="or-key")
|
||||
result = provider.chat_completion(
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
model="anthropic/claude-3.5-sonnet",
|
||||
)
|
||||
|
||||
assert result == "Router response"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# LiteLLMProvider
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestLiteLLMProvider:
|
||||
"""Tests for LiteLLMProvider."""
|
||||
|
||||
@patch("litellm.completion")
|
||||
def test_chat_completion_basic(self, mock_completion):
|
||||
"""Provider calls litellm.completion with correct arguments."""
|
||||
mock_completion.return_value = _make_litellm_response("LiteLLM response")
|
||||
|
||||
provider = LiteLLMProvider()
|
||||
result = provider.chat_completion(
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
model="openai/gpt-4o",
|
||||
)
|
||||
|
||||
assert result == "LiteLLM response"
|
||||
call_kwargs = mock_completion.call_args[1]
|
||||
assert call_kwargs["model"] == "openai/gpt-4o"
|
||||
assert call_kwargs["temperature"] == 0
|
||||
|
||||
@patch("litellm.completion")
|
||||
def test_chat_completion_with_api_key(self, mock_completion):
|
||||
"""Provider forwards api_key when set."""
|
||||
mock_completion.return_value = _make_litellm_response("ok")
|
||||
|
||||
provider = LiteLLMProvider(api_key="my-key")
|
||||
provider.chat_completion(
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
model="openai/gpt-4o",
|
||||
)
|
||||
|
||||
call_kwargs = mock_completion.call_args[1]
|
||||
assert call_kwargs["api_key"] == "my-key"
|
||||
|
||||
@patch("litellm.completion")
|
||||
def test_chat_completion_with_api_base(self, mock_completion):
|
||||
"""Provider forwards api_base when set."""
|
||||
mock_completion.return_value = _make_litellm_response("ok")
|
||||
|
||||
provider = LiteLLMProvider(api_base="http://my-proxy/v1")
|
||||
provider.chat_completion(
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
model="openai/gpt-4o",
|
||||
)
|
||||
|
||||
call_kwargs = mock_completion.call_args[1]
|
||||
assert call_kwargs["api_base"] == "http://my-proxy/v1"
|
||||
|
||||
@patch("litellm.completion")
|
||||
def test_chat_completion_omits_none_api_key(self, mock_completion):
|
||||
"""Provider does not include api_key when it is None."""
|
||||
mock_completion.return_value = _make_litellm_response("ok")
|
||||
|
||||
provider = LiteLLMProvider() # api_key=None
|
||||
provider.chat_completion(
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
model="openai/gpt-4o",
|
||||
)
|
||||
|
||||
call_kwargs = mock_completion.call_args[1]
|
||||
assert "api_key" not in call_kwargs
|
||||
|
||||
@patch("litellm.completion")
|
||||
def test_chat_completion_omits_none_api_base(self, mock_completion):
|
||||
"""Provider does not include api_base when it is None."""
|
||||
mock_completion.return_value = _make_litellm_response("ok")
|
||||
|
||||
provider = LiteLLMProvider() # api_base=None
|
||||
provider.chat_completion(
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
model="openai/gpt-4o",
|
||||
)
|
||||
|
||||
call_kwargs = mock_completion.call_args[1]
|
||||
assert "api_base" not in call_kwargs
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# get_ai_provider factory function
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestGetAIProvider:
|
||||
"""Tests for the get_ai_provider factory function."""
|
||||
|
||||
def _mock_settings(self, **kwargs):
|
||||
"""Return a mock settings object with sensible defaults."""
|
||||
defaults = {
|
||||
"ai_provider": "openai",
|
||||
"openai_api_key": "sk-test",
|
||||
"openai_base_url": "https://api.openai.com/v1",
|
||||
"azure_openai_api_version": "2024-02-01",
|
||||
"anthropic_api_key": None,
|
||||
"gemini_api_key": None,
|
||||
"ollama_base_url": "http://localhost:11434",
|
||||
"openrouter_api_key": None,
|
||||
"openrouter_base_url": "https://openrouter.ai/api/v1",
|
||||
}
|
||||
defaults.update(kwargs)
|
||||
mock_settings = MagicMock()
|
||||
for key, value in defaults.items():
|
||||
setattr(mock_settings, key, value)
|
||||
return mock_settings
|
||||
|
||||
@patch("openai.OpenAI")
|
||||
def test_returns_openai_provider_by_default(self, mock_openai_cls):
|
||||
"""get_ai_provider returns an OpenAIProvider for ai_provider='openai'."""
|
||||
with patch("app.utils.ai_provider.settings", self._mock_settings(ai_provider="openai")):
|
||||
provider = get_ai_provider()
|
||||
assert isinstance(provider, OpenAIProvider)
|
||||
|
||||
@patch("openai.AzureOpenAI")
|
||||
def test_returns_azure_provider(self, mock_azure_cls):
|
||||
"""get_ai_provider returns AzureOpenAIProvider for ai_provider='azure'."""
|
||||
with patch(
|
||||
"app.utils.ai_provider.settings",
|
||||
self._mock_settings(
|
||||
ai_provider="azure",
|
||||
openai_base_url="https://my-resource.openai.azure.com",
|
||||
),
|
||||
):
|
||||
provider = get_ai_provider()
|
||||
assert isinstance(provider, AzureOpenAIProvider)
|
||||
|
||||
def test_returns_anthropic_provider(self):
|
||||
"""get_ai_provider returns AnthropicProvider for ai_provider='anthropic'."""
|
||||
with patch(
|
||||
"app.utils.ai_provider.settings",
|
||||
self._mock_settings(ai_provider="anthropic", anthropic_api_key="ant-key"),
|
||||
):
|
||||
provider = get_ai_provider()
|
||||
assert isinstance(provider, AnthropicProvider)
|
||||
|
||||
def test_anthropic_raises_without_api_key(self):
|
||||
"""get_ai_provider raises ValueError when anthropic_api_key is missing."""
|
||||
with patch(
|
||||
"app.utils.ai_provider.settings",
|
||||
self._mock_settings(ai_provider="anthropic", anthropic_api_key=None),
|
||||
):
|
||||
with pytest.raises(ValueError, match="ANTHROPIC_API_KEY"):
|
||||
get_ai_provider()
|
||||
|
||||
def test_returns_gemini_provider(self):
|
||||
"""get_ai_provider returns GeminiProvider for ai_provider='gemini'."""
|
||||
with patch(
|
||||
"app.utils.ai_provider.settings",
|
||||
self._mock_settings(ai_provider="gemini", gemini_api_key="gemini-key"),
|
||||
):
|
||||
provider = get_ai_provider()
|
||||
assert isinstance(provider, GeminiProvider)
|
||||
|
||||
def test_gemini_raises_without_api_key(self):
|
||||
"""get_ai_provider raises ValueError when gemini_api_key is missing."""
|
||||
with patch(
|
||||
"app.utils.ai_provider.settings",
|
||||
self._mock_settings(ai_provider="gemini", gemini_api_key=None),
|
||||
):
|
||||
with pytest.raises(ValueError, match="GEMINI_API_KEY"):
|
||||
get_ai_provider()
|
||||
|
||||
@patch("openai.OpenAI")
|
||||
def test_returns_ollama_provider(self, mock_openai_cls):
|
||||
"""get_ai_provider returns OllamaProvider for ai_provider='ollama'."""
|
||||
with patch(
|
||||
"app.utils.ai_provider.settings",
|
||||
self._mock_settings(ai_provider="ollama"),
|
||||
):
|
||||
provider = get_ai_provider()
|
||||
assert isinstance(provider, OllamaProvider)
|
||||
|
||||
@patch("openai.OpenAI")
|
||||
def test_returns_openrouter_provider(self, mock_openai_cls):
|
||||
"""get_ai_provider returns OpenRouterProvider for ai_provider='openrouter'."""
|
||||
with patch(
|
||||
"app.utils.ai_provider.settings",
|
||||
self._mock_settings(ai_provider="openrouter", openrouter_api_key="or-key"),
|
||||
):
|
||||
provider = get_ai_provider()
|
||||
assert isinstance(provider, OpenRouterProvider)
|
||||
|
||||
def test_openrouter_raises_without_api_key(self):
|
||||
"""get_ai_provider raises ValueError when openrouter_api_key is missing."""
|
||||
with patch(
|
||||
"app.utils.ai_provider.settings",
|
||||
self._mock_settings(ai_provider="openrouter", openrouter_api_key=None),
|
||||
):
|
||||
with pytest.raises(ValueError, match="OPENROUTER_API_KEY"):
|
||||
get_ai_provider()
|
||||
|
||||
def test_returns_litellm_provider(self):
|
||||
"""get_ai_provider returns LiteLLMProvider for ai_provider='litellm'."""
|
||||
with patch(
|
||||
"app.utils.ai_provider.settings",
|
||||
self._mock_settings(ai_provider="litellm"),
|
||||
):
|
||||
provider = get_ai_provider()
|
||||
assert isinstance(provider, LiteLLMProvider)
|
||||
|
||||
def test_raises_for_unknown_provider(self):
|
||||
"""get_ai_provider raises ValueError for an unrecognised provider name."""
|
||||
with patch(
|
||||
"app.utils.ai_provider.settings",
|
||||
self._mock_settings(ai_provider="unsupported_provider"),
|
||||
):
|
||||
with pytest.raises(ValueError, match="Unknown AI provider"):
|
||||
get_ai_provider()
|
||||
|
||||
def test_provider_name_is_case_insensitive(self):
|
||||
"""get_ai_provider normalises provider names to lowercase."""
|
||||
with patch(
|
||||
"app.utils.ai_provider.settings",
|
||||
self._mock_settings(ai_provider="anthropic", anthropic_api_key="ant-key"),
|
||||
):
|
||||
# "Anthropic" and "ANTHROPIC" should work just like "anthropic"
|
||||
provider = get_ai_provider()
|
||||
assert isinstance(provider, AnthropicProvider)
|
||||
|
||||
def test_litellm_passes_non_default_base_url(self):
|
||||
"""LiteLLM provider receives api_base when openai_base_url differs from default."""
|
||||
with patch(
|
||||
"app.utils.ai_provider.settings",
|
||||
self._mock_settings(
|
||||
ai_provider="litellm",
|
||||
openai_base_url="http://my-proxy/v1",
|
||||
),
|
||||
):
|
||||
provider = get_ai_provider()
|
||||
assert isinstance(provider, LiteLLMProvider)
|
||||
assert provider._api_base == "http://my-proxy/v1"
|
||||
|
||||
def test_litellm_omits_base_url_when_default(self):
|
||||
"""LiteLLM provider has api_base=None when openai_base_url is the default."""
|
||||
with patch(
|
||||
"app.utils.ai_provider.settings",
|
||||
self._mock_settings(
|
||||
ai_provider="litellm",
|
||||
openai_base_url="https://api.openai.com/v1",
|
||||
),
|
||||
):
|
||||
provider = get_ai_provider()
|
||||
assert isinstance(provider, LiteLLMProvider)
|
||||
assert provider._api_base is None
|
||||
@@ -67,12 +67,11 @@ class TestExtractMetadataWithGpt:
|
||||
|
||||
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.client")
|
||||
def test_successful_metadata_extraction(self, mock_client, mock_log_progress, mock_embed_task):
|
||||
"""Test successful metadata extraction with valid GPT response."""
|
||||
# Mock the OpenAI client response
|
||||
mock_completion = MagicMock()
|
||||
mock_completion.choices[0].message.content = json.dumps(
|
||||
@patch("app.utils.ai_provider.get_ai_provider")
|
||||
def test_successful_metadata_extraction(self, mock_get_provider, mock_log_progress, mock_embed_task):
|
||||
"""Test successful metadata extraction with valid AI provider response."""
|
||||
mock_provider = MagicMock()
|
||||
mock_provider.chat_completion.return_value = json.dumps(
|
||||
{
|
||||
"filename": "2024-01-15_Invoice_Amazon",
|
||||
"empfaenger": "John Doe",
|
||||
@@ -89,7 +88,7 @@ class TestExtractMetadataWithGpt:
|
||||
"monetary_amounts": ["99.99 EUR"],
|
||||
}
|
||||
)
|
||||
mock_client.chat.completions.create.return_value = mock_completion
|
||||
mock_get_provider.return_value = mock_provider
|
||||
|
||||
# Set task request context directly on the Celery task
|
||||
extract_metadata_with_gpt.request.id = "test-task-id"
|
||||
@@ -97,11 +96,11 @@ class TestExtractMetadataWithGpt:
|
||||
# Call the underlying function directly (not through Celery)
|
||||
result = extract_metadata_with_gpt.__wrapped__("test_invoice.pdf", "Invoice from Amazon for 99.99 EUR", 123)
|
||||
|
||||
# Verify OpenAI was called
|
||||
mock_client.chat.completions.create.assert_called_once()
|
||||
call_args = mock_client.chat.completions.create.call_args
|
||||
assert call_args[1]["temperature"] == 0
|
||||
assert len(call_args[1]["messages"]) == 2
|
||||
# Verify AI provider was called
|
||||
mock_provider.chat_completion.assert_called_once()
|
||||
call_kwargs = mock_provider.chat_completion.call_args[1]
|
||||
assert call_kwargs["temperature"] == 0
|
||||
assert len(call_kwargs["messages"]) == 2
|
||||
|
||||
# Verify metadata was extracted correctly
|
||||
assert result["s3_file"] == "test_invoice.pdf"
|
||||
@@ -117,14 +116,14 @@ class TestExtractMetadataWithGpt:
|
||||
|
||||
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.client")
|
||||
def test_handles_json_in_backticks(self, mock_client, mock_log_progress, mock_embed_task):
|
||||
@patch("app.utils.ai_provider.get_ai_provider")
|
||||
def test_handles_json_in_backticks(self, mock_get_provider, mock_log_progress, mock_embed_task):
|
||||
"""Test extraction handles JSON wrapped in markdown code blocks."""
|
||||
mock_completion = MagicMock()
|
||||
mock_completion.choices[
|
||||
0
|
||||
].message.content = '```json\n{"filename": "test.pdf", "document_type": "Unknown"}\n```'
|
||||
mock_client.chat.completions.create.return_value = mock_completion
|
||||
mock_provider = MagicMock()
|
||||
mock_provider.chat_completion.return_value = (
|
||||
'```json\n{"filename": "test.pdf", "document_type": "Unknown"}\n```'
|
||||
)
|
||||
mock_get_provider.return_value = mock_provider
|
||||
|
||||
extract_metadata_with_gpt.request.id = "test-task-id"
|
||||
|
||||
@@ -136,12 +135,12 @@ class TestExtractMetadataWithGpt:
|
||||
|
||||
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.client")
|
||||
def test_handles_invalid_json_response(self, mock_client, mock_log_progress, mock_embed_task):
|
||||
"""Test handling of invalid JSON in GPT response."""
|
||||
mock_completion = MagicMock()
|
||||
mock_completion.choices[0].message.content = "This is not valid JSON at all"
|
||||
mock_client.chat.completions.create.return_value = mock_completion
|
||||
@patch("app.utils.ai_provider.get_ai_provider")
|
||||
def test_handles_invalid_json_response(self, mock_get_provider, mock_log_progress, mock_embed_task):
|
||||
"""Test handling of invalid JSON in AI provider response."""
|
||||
mock_provider = MagicMock()
|
||||
mock_provider.chat_completion.return_value = "This is not valid JSON at all"
|
||||
mock_get_provider.return_value = mock_provider
|
||||
|
||||
extract_metadata_with_gpt.request.id = "test-task-id"
|
||||
|
||||
@@ -155,10 +154,12 @@ class TestExtractMetadataWithGpt:
|
||||
|
||||
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.client")
|
||||
def test_handles_openai_api_exception(self, mock_client, mock_log_progress, mock_embed_task):
|
||||
"""Test handling of OpenAI API exceptions."""
|
||||
mock_client.chat.completions.create.side_effect = Exception("API Error: Rate limit exceeded")
|
||||
@patch("app.utils.ai_provider.get_ai_provider")
|
||||
def test_handles_api_exception(self, mock_get_provider, mock_log_progress, mock_embed_task):
|
||||
"""Test handling of AI provider API exceptions."""
|
||||
mock_provider = MagicMock()
|
||||
mock_provider.chat_completion.side_effect = Exception("API Error: Rate limit exceeded")
|
||||
mock_get_provider.return_value = mock_provider
|
||||
|
||||
extract_metadata_with_gpt.request.id = "test-task-id"
|
||||
|
||||
@@ -172,15 +173,15 @@ class TestExtractMetadataWithGpt:
|
||||
|
||||
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.client")
|
||||
@patch("app.utils.ai_provider.get_ai_provider")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.SessionLocal")
|
||||
def test_retrieves_file_id_from_database_when_not_provided(
|
||||
self, mock_session_local, mock_client, mock_log_progress, mock_embed_task
|
||||
self, mock_session_local, mock_get_provider, mock_log_progress, mock_embed_task
|
||||
):
|
||||
"""Test file_id retrieval from database when not provided."""
|
||||
mock_completion = MagicMock()
|
||||
mock_completion.choices[0].message.content = '{"filename": "test.pdf", "document_type": "Unknown"}'
|
||||
mock_client.chat.completions.create.return_value = mock_completion
|
||||
mock_provider = MagicMock()
|
||||
mock_provider.chat_completion.return_value = '{"filename": "test.pdf", "document_type": "Unknown"}'
|
||||
mock_get_provider.return_value = mock_provider
|
||||
|
||||
# Mock database session
|
||||
mock_db = MagicMock()
|
||||
@@ -206,15 +207,15 @@ class TestExtractMetadataWithGpt:
|
||||
|
||||
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.client")
|
||||
def test_validates_filename_security(self, mock_client, mock_log_progress, mock_embed_task):
|
||||
@patch("app.utils.ai_provider.get_ai_provider")
|
||||
def test_validates_filename_security(self, mock_get_provider, mock_log_progress, mock_embed_task):
|
||||
"""Test filename validation to prevent path traversal."""
|
||||
mock_completion = MagicMock()
|
||||
mock_provider = MagicMock()
|
||||
# Try to inject a malicious filename
|
||||
mock_completion.choices[0].message.content = json.dumps(
|
||||
mock_provider.chat_completion.return_value = json.dumps(
|
||||
{"filename": "../../../etc/passwd", "document_type": "Invoice"}
|
||||
)
|
||||
mock_client.chat.completions.create.return_value = mock_completion
|
||||
mock_get_provider.return_value = mock_provider
|
||||
|
||||
extract_metadata_with_gpt.request.id = "test-task-id"
|
||||
|
||||
@@ -226,14 +227,14 @@ class TestExtractMetadataWithGpt:
|
||||
|
||||
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.client")
|
||||
def test_validates_filename_with_dots(self, mock_client, mock_log_progress, mock_embed_task):
|
||||
@patch("app.utils.ai_provider.get_ai_provider")
|
||||
def test_validates_filename_with_dots(self, mock_get_provider, mock_log_progress, mock_embed_task):
|
||||
"""Test filename validation rejects '..' in filenames."""
|
||||
mock_completion = MagicMock()
|
||||
mock_completion.choices[0].message.content = json.dumps(
|
||||
mock_provider = MagicMock()
|
||||
mock_provider.chat_completion.return_value = json.dumps(
|
||||
{"filename": "test..invoice.pdf", "document_type": "Invoice"}
|
||||
)
|
||||
mock_client.chat.completions.create.return_value = mock_completion
|
||||
mock_get_provider.return_value = mock_provider
|
||||
|
||||
extract_metadata_with_gpt.request.id = "test-task-id"
|
||||
|
||||
@@ -244,14 +245,14 @@ class TestExtractMetadataWithGpt:
|
||||
|
||||
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.client")
|
||||
def test_accepts_valid_filename(self, mock_client, mock_log_progress, mock_embed_task):
|
||||
@patch("app.utils.ai_provider.get_ai_provider")
|
||||
def test_accepts_valid_filename(self, mock_get_provider, mock_log_progress, mock_embed_task):
|
||||
"""Test that valid filenames are accepted."""
|
||||
mock_completion = MagicMock()
|
||||
mock_completion.choices[0].message.content = json.dumps(
|
||||
mock_provider = MagicMock()
|
||||
mock_provider.chat_completion.return_value = json.dumps(
|
||||
{"filename": "2024-01-15_Invoice_Amazon.pdf", "document_type": "Invoice"}
|
||||
)
|
||||
mock_client.chat.completions.create.return_value = mock_completion
|
||||
mock_get_provider.return_value = mock_provider
|
||||
|
||||
extract_metadata_with_gpt.request.id = "test-task-id"
|
||||
|
||||
@@ -262,12 +263,12 @@ class TestExtractMetadataWithGpt:
|
||||
|
||||
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.client")
|
||||
def test_handles_malformed_json_with_valid_structure(self, mock_client, mock_log_progress, mock_embed_task):
|
||||
@patch("app.utils.ai_provider.get_ai_provider")
|
||||
def test_handles_malformed_json_with_valid_structure(self, mock_get_provider, mock_log_progress, mock_embed_task):
|
||||
"""Test handling of JSON that's parseable but missing expected fields."""
|
||||
mock_completion = MagicMock()
|
||||
mock_completion.choices[0].message.content = '{"unexpected_field": "value"}'
|
||||
mock_client.chat.completions.create.return_value = mock_completion
|
||||
mock_provider = MagicMock()
|
||||
mock_provider.chat_completion.return_value = '{"unexpected_field": "value"}'
|
||||
mock_get_provider.return_value = mock_provider
|
||||
|
||||
extract_metadata_with_gpt.request.id = "test-task-id"
|
||||
|
||||
@@ -279,12 +280,12 @@ class TestExtractMetadataWithGpt:
|
||||
|
||||
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.client")
|
||||
def test_handles_absolute_path_filename(self, mock_client, mock_log_progress, mock_embed_task):
|
||||
"""Test handling when filename is provided as an absolute path (line 73)."""
|
||||
mock_completion = MagicMock()
|
||||
mock_completion.choices[0].message.content = '{"filename": "test.pdf", "document_type": "Unknown"}'
|
||||
mock_client.chat.completions.create.return_value = mock_completion
|
||||
@patch("app.utils.ai_provider.get_ai_provider")
|
||||
def test_handles_absolute_path_filename(self, mock_get_provider, mock_log_progress, mock_embed_task):
|
||||
"""Test handling when filename is provided as an absolute path."""
|
||||
mock_provider = MagicMock()
|
||||
mock_provider.chat_completion.return_value = '{"filename": "test.pdf", "document_type": "Unknown"}'
|
||||
mock_get_provider.return_value = mock_provider
|
||||
|
||||
extract_metadata_with_gpt.request.id = "test-task-id"
|
||||
|
||||
@@ -298,15 +299,15 @@ class TestExtractMetadataWithGpt:
|
||||
|
||||
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.client")
|
||||
@patch("app.utils.ai_provider.get_ai_provider")
|
||||
@patch("app.tasks.extract_metadata_with_gpt.SessionLocal")
|
||||
def test_database_lookup_with_existing_file(
|
||||
self, mock_session_local, mock_client, mock_log_progress, mock_embed_task
|
||||
self, mock_session_local, mock_get_provider, mock_log_progress, mock_embed_task
|
||||
):
|
||||
"""Test file_id retrieval when file exists on disk and in database (branches 76->82, 79->82)."""
|
||||
mock_completion = MagicMock()
|
||||
mock_completion.choices[0].message.content = '{"filename": "test.pdf", "document_type": "Unknown"}'
|
||||
mock_client.chat.completions.create.return_value = mock_completion
|
||||
"""Test file_id retrieval when file exists on disk and in database."""
|
||||
mock_provider = MagicMock()
|
||||
mock_provider.chat_completion.return_value = '{"filename": "test.pdf", "document_type": "Unknown"}'
|
||||
mock_get_provider.return_value = mock_provider
|
||||
|
||||
# Mock database session
|
||||
mock_db = MagicMock()
|
||||
@@ -333,20 +334,14 @@ class TestExtractMetadataWithGpt:
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestClientInitialization:
|
||||
"""Tests for OpenAI client initialization error handling."""
|
||||
class TestModuleImports:
|
||||
"""Tests for module import behaviour."""
|
||||
|
||||
def test_client_initialization_imports_successfully(self):
|
||||
"""Test that module imports successfully even if client initialization fails (lines 25-27).
|
||||
def test_module_imports_successfully(self):
|
||||
"""Test that the module imports successfully."""
|
||||
import app.tasks.extract_metadata_with_gpt as mod
|
||||
|
||||
The module has a try/except block for client initialization that sets client to None
|
||||
on failure. This test verifies the module can be imported without crashing,
|
||||
regardless of whether the client initializes successfully or not.
|
||||
"""
|
||||
# Import should succeed regardless of client initialization success
|
||||
from app.tasks.extract_metadata_with_gpt import client
|
||||
|
||||
# Client will be either an OpenAI client instance or None
|
||||
# Both are valid states - the important thing is the import doesn't crash
|
||||
# We verify the client variable exists and has a defined type
|
||||
assert hasattr(client, "__class__") or client is None
|
||||
assert mod is not None
|
||||
assert hasattr(mod, "extract_metadata_with_gpt")
|
||||
assert hasattr(mod, "extract_json_from_text")
|
||||
assert hasattr(mod, "get_ai_provider")
|
||||
|
||||
@@ -407,65 +407,60 @@ startxref
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestRefineTextWithGPT:
|
||||
"""Tests for OpenAI text refinement task."""
|
||||
"""Tests for AI provider text refinement task."""
|
||||
|
||||
@patch("app.tasks.refine_text_with_gpt.log_task_progress")
|
||||
def test_successful_text_refinement(self, mock_log):
|
||||
"""Test successful text refinement with OpenAI."""
|
||||
"""Test successful text refinement with AI provider."""
|
||||
raw_text = "This is s0me text with OCR err0rs"
|
||||
filename = "test.pdf"
|
||||
|
||||
# Mock OpenAI response
|
||||
mock_choice = Mock()
|
||||
mock_choice.message.content = "This is some text with OCR errors"
|
||||
|
||||
mock_response = Mock()
|
||||
mock_response.choices = [mock_choice]
|
||||
|
||||
mock_client = Mock()
|
||||
mock_client.chat.completions.create.return_value = mock_response
|
||||
cleaned = "This is some text with OCR errors"
|
||||
|
||||
# Import the module to patch the correct function
|
||||
from app.tasks import extract_metadata_with_gpt as metadata_module
|
||||
|
||||
mock_provider = MagicMock()
|
||||
mock_provider.chat_completion.return_value = cleaned
|
||||
|
||||
with (
|
||||
patch("app.tasks.refine_text_with_gpt.client", mock_client),
|
||||
patch("app.utils.ai_provider.get_ai_provider", return_value=mock_provider),
|
||||
patch.object(metadata_module, "extract_metadata_with_gpt") as mock_extract,
|
||||
patch("app.tasks.refine_text_with_gpt.settings") as mock_settings,
|
||||
):
|
||||
mock_settings.openai_model = "gpt-4"
|
||||
mock_settings.ai_model = None
|
||||
mock_extract.delay = MagicMock()
|
||||
|
||||
result = refine_text_with_gpt.run(filename, raw_text)
|
||||
|
||||
# Verify results
|
||||
assert result["filename"] == filename
|
||||
assert result["cleaned_text"] == "This is some text with OCR errors"
|
||||
assert result["cleaned_text"] == cleaned
|
||||
|
||||
# Verify OpenAI was called correctly
|
||||
mock_client.chat.completions.create.assert_called_once()
|
||||
call_kwargs = mock_client.chat.completions.create.call_args[1]
|
||||
assert call_kwargs["model"] == "gpt-4"
|
||||
# Verify AI provider was called correctly
|
||||
mock_provider.chat_completion.assert_called_once()
|
||||
call_kwargs = mock_provider.chat_completion.call_args[1]
|
||||
assert len(call_kwargs["messages"]) == 2
|
||||
assert call_kwargs["messages"][1]["content"] == raw_text
|
||||
|
||||
# Verify metadata extraction was queued
|
||||
mock_extract.delay.assert_called_once_with(filename, "This is some text with OCR errors")
|
||||
mock_extract.delay.assert_called_once_with(filename, cleaned)
|
||||
|
||||
@patch("app.tasks.refine_text_with_gpt.log_task_progress")
|
||||
def test_openai_api_error(self, mock_log):
|
||||
"""Test error handling when OpenAI API fails."""
|
||||
"""Test error handling when AI provider call fails."""
|
||||
raw_text = "Test text"
|
||||
filename = "test.pdf"
|
||||
|
||||
mock_client = Mock()
|
||||
mock_client.chat.completions.create.side_effect = Exception("OpenAI API error")
|
||||
mock_provider = MagicMock()
|
||||
mock_provider.chat_completion.side_effect = Exception("OpenAI API error")
|
||||
|
||||
with (
|
||||
patch("app.tasks.refine_text_with_gpt.client", mock_client),
|
||||
patch("app.utils.ai_provider.get_ai_provider", return_value=mock_provider),
|
||||
patch("app.tasks.refine_text_with_gpt.settings") as mock_settings,
|
||||
):
|
||||
mock_settings.openai_model = "gpt-4"
|
||||
mock_settings.ai_model = None
|
||||
|
||||
# Should raise the exception
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
|
||||
Reference in New Issue
Block a user