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:
copilot-swe-agent[bot]
2026-02-23 19:26:45 +00:00
parent 0736cd8710
commit d4c7fb26ac
8 changed files with 1075 additions and 136 deletions
+591
View File
@@ -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
+76 -81
View File
@@ -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")
+18 -23
View File
@@ -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: