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
+22
View File
@@ -15,6 +15,28 @@ class Settings(BaseSettings):
openai_api_key: str openai_api_key: str
openai_base_url: str = "https://api.openai.com/v1" # Default to OpenAI's endpoint openai_base_url: str = "https://api.openai.com/v1" # Default to OpenAI's endpoint
openai_model: str = "gpt-4o-mini" # Default model openai_model: str = "gpt-4o-mini" # Default model
# AI provider abstraction layer
# Supported values: openai, azure, anthropic, gemini, ollama, openrouter, litellm
ai_provider: str = "openai"
# Override model for any provider; falls back to openai_model when not set
ai_model: Optional[str] = None
# Anthropic Claude settings (used when ai_provider="anthropic")
anthropic_api_key: Optional[str] = None
# Google Gemini settings (used when ai_provider="gemini")
gemini_api_key: Optional[str] = None
# Ollama local LLM settings (used when ai_provider="ollama")
ollama_base_url: str = "http://localhost:11434"
# OpenRouter settings (used when ai_provider="openrouter")
openrouter_api_key: Optional[str] = None
openrouter_base_url: str = "https://openrouter.ai/api/v1"
# Azure OpenAI API version (used when ai_provider="azure")
azure_openai_api_version: str = "2024-02-01"
workdir: str workdir: str
debug: bool = False # Default to False debug: bool = False # Default to False
+10 -18
View File
@@ -5,8 +5,6 @@ import logging
import os import os
import re import re
import openai
# Import the shared Celery instance # Import the shared Celery instance
from app.celery_app import celery from app.celery_app import celery
from app.config import settings from app.config import settings
@@ -15,17 +13,10 @@ from app.models import FileRecord
from app.tasks.embed_metadata_into_pdf import embed_metadata_into_pdf from app.tasks.embed_metadata_into_pdf import embed_metadata_into_pdf
from app.tasks.retry_config import BaseTaskWithRetry from app.tasks.retry_config import BaseTaskWithRetry
from app.utils import log_task_progress from app.utils import log_task_progress
from app.utils.ai_provider import get_ai_provider
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# Initialize OpenAI client dynamically with better error handling
try:
client = openai.OpenAI(api_key=settings.openai_api_key, base_url=settings.openai_base_url)
logger.info("OpenAI client initialized successfully")
except Exception as e:
logger.error(f"Failed to initialize OpenAI client: {e}")
client = None
def extract_json_from_text(text): def extract_json_from_text(text):
""" """
@@ -114,23 +105,24 @@ def extract_metadata_with_gpt(self, filename: str, cleaned_text: str, file_id: i
try: try:
logger.info(f"[{task_id}] Sending classification request for {filename}...") logger.info(f"[{task_id}] Sending classification request for {filename}...")
log_task_progress(task_id, "call_openai", "in_progress", "Calling OpenAI API", file_id=file_id) log_task_progress(task_id, "call_ai_provider", "in_progress", "Calling AI provider API", file_id=file_id)
completion = client.chat.completions.create( provider = get_ai_provider()
model=settings.openai_model, model = settings.ai_model or settings.openai_model
content = provider.chat_completion(
messages=[ messages=[
{"role": "system", "content": "You are an intelligent document classifier."}, {"role": "system", "content": "You are an intelligent document classifier."},
{"role": "user", "content": prompt}, {"role": "user", "content": prompt},
], ],
model=model,
temperature=0, temperature=0,
) )
content = completion.choices[0].message.content
logger.info(f"[{task_id}] Raw classification response for {filename}: {content[:200]}...") logger.info(f"[{task_id}] Raw classification response for {filename}: {content[:200]}...")
log_task_progress( log_task_progress(
task_id, task_id,
"call_openai", "call_ai_provider",
"success", "success",
"Received OpenAI response", "Received AI provider response",
file_id=file_id, file_id=file_id,
detail=f"Raw classification response:\n{content}", detail=f"Raw classification response:\n{content}",
) )
@@ -187,13 +179,13 @@ def extract_metadata_with_gpt(self, filename: str, cleaned_text: str, file_id: i
return {"s3_file": os.path.basename(filename), "metadata": metadata} return {"s3_file": os.path.basename(filename), "metadata": metadata}
except Exception as e: except Exception as e:
logger.exception(f"[{task_id}] OpenAI classification failed for {filename}: {e}") logger.exception(f"[{task_id}] AI provider classification failed for {filename}: {e}")
log_task_progress( log_task_progress(
task_id, task_id,
"extract_metadata_with_gpt", "extract_metadata_with_gpt",
"failure", "failure",
f"Exception: {str(e)}", f"Exception: {str(e)}",
file_id=file_id, file_id=file_id,
detail=f"OpenAI classification failed for {filename}.\nException: {str(e)}", detail=f"AI provider classification failed for {filename}.\nException: {str(e)}",
) )
return {} return {}
+9 -13
View File
@@ -2,32 +2,29 @@
import logging import logging
import openai
# Import the shared Celery instance # Import the shared Celery instance
from app.celery_app import celery from app.celery_app import celery
from app.config import settings from app.config import settings
from app.tasks.retry_config import BaseTaskWithRetry from app.tasks.retry_config import BaseTaskWithRetry
from app.utils import log_task_progress from app.utils import log_task_progress
from app.utils.ai_provider import get_ai_provider
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# Initialize OpenAI client dynamically
client = openai.OpenAI(api_key=settings.openai_api_key, base_url=settings.openai_base_url)
@celery.task(base=BaseTaskWithRetry, bind=True) @celery.task(base=BaseTaskWithRetry, bind=True)
def refine_text_with_gpt(self, filename: str, raw_text: str): def refine_text_with_gpt(self, filename: str, raw_text: str):
"""Uses OpenAI to clean and refine OCR text.""" """Uses the configured AI provider to clean and refine OCR text."""
task_id = self.request.id task_id = self.request.id
logger.info(f"[{task_id}] Starting OCR text refinement for: {filename}") logger.info(f"[{task_id}] Starting OCR text refinement for: {filename}")
log_task_progress(task_id, "refine_text_with_gpt", "in_progress", f"Refining OCR text for {filename}") log_task_progress(task_id, "refine_text_with_gpt", "in_progress", f"Refining OCR text for {filename}")
try: try:
log_task_progress(task_id, "call_openai", "in_progress", "Calling OpenAI for text refinement") log_task_progress(task_id, "call_ai_provider", "in_progress", "Calling AI provider for text refinement")
response = client.chat.completions.create( provider = get_ai_provider()
model=settings.openai_model, model = settings.ai_model or settings.openai_model
cleaned_text = provider.chat_completion(
messages=[ messages=[
{ {
"role": "system", "role": "system",
@@ -38,16 +35,15 @@ def refine_text_with_gpt(self, filename: str, raw_text: str):
}, },
{"role": "user", "content": raw_text}, {"role": "user", "content": raw_text},
], ],
model=model,
) )
cleaned_text = response.choices[0].message.content
logger.info(f"[{task_id}] Text refinement complete for {filename}: {len(cleaned_text)} characters") logger.info(f"[{task_id}] Text refinement complete for {filename}: {len(cleaned_text)} characters")
log_task_progress( log_task_progress(
task_id, task_id,
"call_openai", "call_ai_provider",
"success", "success",
"Received refined text from OpenAI", "Received refined text from AI provider",
detail=f"Input: {len(raw_text)} chars → Output: {len(cleaned_text)} chars", detail=f"Input: {len(raw_text)} chars → Output: {len(cleaned_text)} chars",
) )
+345
View File
@@ -0,0 +1,345 @@
#!/usr/bin/env python3
"""AI provider abstraction layer for DocuElevate.
This module provides a pluggable abstraction for various AI model providers,
allowing the platform to work with OpenAI, Azure OpenAI, Anthropic Claude,
Google Gemini, Ollama (local LLMs), OpenRouter, and any LiteLLM-compatible
provider without being locked to a single vendor.
Provider selection is controlled by the ``AI_PROVIDER`` environment variable.
See the Configuration Guide for full details on each provider's settings.
"""
import logging
from abc import ABC, abstractmethod
from typing import Any, Dict, List, Optional
logger = logging.getLogger(__name__)
class AIProvider(ABC):
"""Abstract base class for AI chat completion providers.
All concrete providers must implement :meth:`chat_completion`, which
accepts a list of chat messages and returns the model's response as a
plain string. The interface intentionally mirrors the OpenAI Chat
Completions API so that callers need no provider-specific knowledge.
"""
@abstractmethod
def chat_completion(
self,
messages: List[Dict[str, str]],
model: str,
temperature: float = 0,
**kwargs: Any,
) -> str:
"""Get a chat completion from the AI provider.
Args:
messages: List of message dicts with ``role`` and ``content`` keys.
model: Model name/identifier to use (provider-specific format).
temperature: Sampling temperature (01). Default: 0 (deterministic).
**kwargs: Additional provider-specific arguments passed through.
Returns:
The model's response as a plain string.
Raises:
Exception: If the underlying API call fails.
"""
class OpenAIProvider(AIProvider):
"""OpenAI provider using the ``openai`` Python SDK.
Also works as a drop-in for any OpenAI-compatible API endpoint, including
LocalAI and LM Studio. Ollama and OpenRouter have dedicated providers with
sensible defaults, but this provider works for them too when a custom
``base_url`` is supplied.
"""
def __init__(self, api_key: str, base_url: Optional[str] = None) -> None:
import openai
self._client = openai.OpenAI(
api_key=api_key,
base_url=base_url or "https://api.openai.com/v1",
)
def chat_completion(
self,
messages: List[Dict[str, str]],
model: str,
temperature: float = 0,
**kwargs: Any,
) -> str:
completion = self._client.chat.completions.create(
model=model,
messages=messages,
temperature=temperature,
**kwargs,
)
return completion.choices[0].message.content
class AzureOpenAIProvider(AIProvider):
"""Azure OpenAI provider using the ``openai`` Python SDK's Azure client."""
def __init__(self, api_key: str, azure_endpoint: str, api_version: str = "2024-02-01") -> None:
import openai
self._client = openai.AzureOpenAI(
api_key=api_key,
azure_endpoint=azure_endpoint,
api_version=api_version,
)
def chat_completion(
self,
messages: List[Dict[str, str]],
model: str,
temperature: float = 0,
**kwargs: Any,
) -> str:
completion = self._client.chat.completions.create(
model=model,
messages=messages,
temperature=temperature,
**kwargs,
)
return completion.choices[0].message.content
class AnthropicProvider(AIProvider):
"""Anthropic Claude provider routed via LiteLLM.
Requires ``litellm`` to be installed. Model names should be in Anthropic
format (e.g. ``claude-3-5-sonnet-20241022``); the ``anthropic/`` prefix is
added automatically when absent.
"""
def __init__(self, api_key: str) -> None:
self._api_key = api_key
def chat_completion(
self,
messages: List[Dict[str, str]],
model: str,
temperature: float = 0,
**kwargs: Any,
) -> str:
import litellm
model_name = model if model.startswith("anthropic/") else f"anthropic/{model}"
response = litellm.completion(
model=model_name,
messages=messages,
temperature=temperature,
api_key=self._api_key,
**kwargs,
)
return response.choices[0].message.content
class GeminiProvider(AIProvider):
"""Google Gemini provider routed via LiteLLM.
Requires ``litellm`` to be installed. Model names should be in Gemini
format (e.g. ``gemini-1.5-pro``); the ``gemini/`` prefix is added
automatically when absent.
"""
def __init__(self, api_key: str) -> None:
self._api_key = api_key
def chat_completion(
self,
messages: List[Dict[str, str]],
model: str,
temperature: float = 0,
**kwargs: Any,
) -> str:
import litellm
model_name = model if model.startswith("gemini/") else f"gemini/{model}"
response = litellm.completion(
model=model_name,
messages=messages,
temperature=temperature,
api_key=self._api_key,
**kwargs,
)
return response.choices[0].message.content
class OllamaProvider(AIProvider):
"""Ollama local LLM provider via its OpenAI-compatible REST API.
Ollama exposes an OpenAI-compatible endpoint at ``/v1``. Any model
pulled into your Ollama instance (e.g. ``llama3.2``, ``qwen2.5``,
``phi3``) can be used directly by name.
For CPU-only deployments the recommended models are:
* ``llama3.2`` (3B) good balance of speed and quality
* ``qwen2.5`` (3B/7B) excellent at structured JSON output
* ``phi3`` (3.8B) strong reasoning, fast on CPU
See https://ollama.com for installation and model management.
"""
def __init__(self, base_url: str = "http://localhost:11434") -> None:
import openai
self._client = openai.OpenAI(
api_key="ollama", # Ollama does not require a real API key
base_url=f"{base_url.rstrip('/')}/v1",
)
def chat_completion(
self,
messages: List[Dict[str, str]],
model: str,
temperature: float = 0,
**kwargs: Any,
) -> str:
completion = self._client.chat.completions.create(
model=model,
messages=messages,
temperature=temperature,
**kwargs,
)
return completion.choices[0].message.content
class OpenRouterProvider(AIProvider):
"""OpenRouter AI aggregator (https://openrouter.ai).
OpenRouter provides access to 100+ models from OpenAI, Anthropic, Google,
Meta, Mistral, and many others through a single OpenAI-compatible endpoint.
Model names use the ``provider/model`` format (e.g.
``anthropic/claude-3.5-sonnet``, ``google/gemini-pro``).
"""
def __init__(self, api_key: str, base_url: str = "https://openrouter.ai/api/v1") -> None:
import openai
self._client = openai.OpenAI(
api_key=api_key,
base_url=base_url,
)
def chat_completion(
self,
messages: List[Dict[str, str]],
model: str,
temperature: float = 0,
**kwargs: Any,
) -> str:
completion = self._client.chat.completions.create(
model=model,
messages=messages,
temperature=temperature,
**kwargs,
)
return completion.choices[0].message.content
class LiteLLMProvider(AIProvider):
"""LiteLLM provider unified interface for 100+ LLMs.
LiteLLM (https://litellm.ai) translates calls to a single interface that
supports OpenAI, Azure, Anthropic, Gemini, Cohere, Ollama, and many more.
Use the LiteLLM model-string format ``provider/model`` (e.g.
``openai/gpt-4o``, ``anthropic/claude-3-5-sonnet-20241022``,
``ollama/llama3.2``).
This provider is useful when you want LiteLLM to handle all routing and
need features like automatic retries, fallbacks, or cost tracking.
"""
def __init__(self, api_key: Optional[str] = None, api_base: Optional[str] = None) -> None:
self._api_key = api_key
self._api_base = api_base
def chat_completion(
self,
messages: List[Dict[str, str]],
model: str,
temperature: float = 0,
**kwargs: Any,
) -> str:
import litellm
completion_kwargs: Dict[str, Any] = {
"model": model,
"messages": messages,
"temperature": temperature,
}
if self._api_key:
completion_kwargs["api_key"] = self._api_key
if self._api_base:
completion_kwargs["api_base"] = self._api_base
completion_kwargs.update(kwargs)
response = litellm.completion(**completion_kwargs)
return response.choices[0].message.content
def get_ai_provider() -> AIProvider:
"""Factory function that creates and returns the configured AI provider.
Reads ``settings.ai_provider`` (set via the ``AI_PROVIDER`` environment
variable) to select the provider implementation. Provider-specific
credentials and URLs are read from their corresponding settings fields.
Returns:
An :class:`AIProvider` instance ready to serve chat completions.
Raises:
ValueError: If the configured provider name is not recognised.
ValueError: If required credentials for the selected provider are absent.
"""
from app.config import settings
provider = settings.ai_provider.lower()
logger.debug(f"Creating AI provider: {provider}")
if provider == "openai":
return OpenAIProvider(
api_key=settings.openai_api_key,
base_url=settings.openai_base_url,
)
elif provider == "azure":
return AzureOpenAIProvider(
api_key=settings.openai_api_key,
azure_endpoint=settings.openai_base_url,
api_version=settings.azure_openai_api_version,
)
elif provider == "anthropic":
if not settings.anthropic_api_key:
raise ValueError("ANTHROPIC_API_KEY must be set when AI_PROVIDER='anthropic'")
return AnthropicProvider(api_key=settings.anthropic_api_key)
elif provider == "gemini":
if not settings.gemini_api_key:
raise ValueError("GEMINI_API_KEY must be set when AI_PROVIDER='gemini'")
return GeminiProvider(api_key=settings.gemini_api_key)
elif provider == "ollama":
return OllamaProvider(base_url=settings.ollama_base_url)
elif provider == "openrouter":
if not settings.openrouter_api_key:
raise ValueError("OPENROUTER_API_KEY must be set when AI_PROVIDER='openrouter'")
return OpenRouterProvider(
api_key=settings.openrouter_api_key,
base_url=settings.openrouter_base_url,
)
elif provider == "litellm":
return LiteLLMProvider(
api_key=settings.openai_api_key or None,
api_base=settings.openai_base_url if settings.openai_base_url != "https://api.openai.com/v1" else None,
)
else:
raise ValueError(
f"Unknown AI provider: '{provider}'. "
"Supported providers: openai, azure, anthropic, gemini, ollama, openrouter, litellm"
)
+3
View File
@@ -35,3 +35,6 @@ paramiko>=3.4.0 # SSH/SFTP implementation for Python (LGPL license)
# Notification service # Notification service
apprise>=1.4.0 apprise>=1.4.0
# AI provider aggregator - enables Anthropic, Gemini, Ollama, and 100+ LLM providers
litellm>=1.0.0,<2.0.0
+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.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @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")
def test_successful_metadata_extraction(self, mock_client, mock_log_progress, mock_embed_task): def test_successful_metadata_extraction(self, mock_get_provider, mock_log_progress, mock_embed_task):
"""Test successful metadata extraction with valid GPT response.""" """Test successful metadata extraction with valid AI provider response."""
# Mock the OpenAI client response mock_provider = MagicMock()
mock_completion = MagicMock() mock_provider.chat_completion.return_value = json.dumps(
mock_completion.choices[0].message.content = json.dumps(
{ {
"filename": "2024-01-15_Invoice_Amazon", "filename": "2024-01-15_Invoice_Amazon",
"empfaenger": "John Doe", "empfaenger": "John Doe",
@@ -89,7 +88,7 @@ class TestExtractMetadataWithGpt:
"monetary_amounts": ["99.99 EUR"], "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 # Set task request context directly on the Celery task
extract_metadata_with_gpt.request.id = "test-task-id" extract_metadata_with_gpt.request.id = "test-task-id"
@@ -97,11 +96,11 @@ class TestExtractMetadataWithGpt:
# Call the underlying function directly (not through Celery) # 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) result = extract_metadata_with_gpt.__wrapped__("test_invoice.pdf", "Invoice from Amazon for 99.99 EUR", 123)
# Verify OpenAI was called # Verify AI provider was called
mock_client.chat.completions.create.assert_called_once() mock_provider.chat_completion.assert_called_once()
call_args = mock_client.chat.completions.create.call_args call_kwargs = mock_provider.chat_completion.call_args[1]
assert call_args[1]["temperature"] == 0 assert call_kwargs["temperature"] == 0
assert len(call_args[1]["messages"]) == 2 assert len(call_kwargs["messages"]) == 2
# Verify metadata was extracted correctly # Verify metadata was extracted correctly
assert result["s3_file"] == "test_invoice.pdf" 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.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @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")
def test_handles_json_in_backticks(self, mock_client, mock_log_progress, mock_embed_task): 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.""" """Test extraction handles JSON wrapped in markdown code blocks."""
mock_completion = MagicMock() mock_provider = MagicMock()
mock_completion.choices[ mock_provider.chat_completion.return_value = (
0 '```json\n{"filename": "test.pdf", "document_type": "Unknown"}\n```'
].message.content = '```json\n{"filename": "test.pdf", "document_type": "Unknown"}\n```' )
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" 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.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @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")
def test_handles_invalid_json_response(self, mock_client, mock_log_progress, mock_embed_task): def test_handles_invalid_json_response(self, mock_get_provider, mock_log_progress, mock_embed_task):
"""Test handling of invalid JSON in GPT response.""" """Test handling of invalid JSON in AI provider response."""
mock_completion = MagicMock() mock_provider = MagicMock()
mock_completion.choices[0].message.content = "This is not valid JSON at all" mock_provider.chat_completion.return_value = "This is not valid JSON at all"
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" 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.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @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")
def test_handles_openai_api_exception(self, mock_client, mock_log_progress, mock_embed_task): def test_handles_api_exception(self, mock_get_provider, mock_log_progress, mock_embed_task):
"""Test handling of OpenAI API exceptions.""" """Test handling of AI provider API exceptions."""
mock_client.chat.completions.create.side_effect = Exception("API Error: Rate limit exceeded") 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" 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.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @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") @patch("app.tasks.extract_metadata_with_gpt.SessionLocal")
def test_retrieves_file_id_from_database_when_not_provided( 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.""" """Test file_id retrieval from database when not provided."""
mock_completion = MagicMock() mock_provider = MagicMock()
mock_completion.choices[0].message.content = '{"filename": "test.pdf", "document_type": "Unknown"}' mock_provider.chat_completion.return_value = '{"filename": "test.pdf", "document_type": "Unknown"}'
mock_client.chat.completions.create.return_value = mock_completion mock_get_provider.return_value = mock_provider
# Mock database session # Mock database session
mock_db = MagicMock() 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.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @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")
def test_validates_filename_security(self, mock_client, mock_log_progress, mock_embed_task): def test_validates_filename_security(self, mock_get_provider, mock_log_progress, mock_embed_task):
"""Test filename validation to prevent path traversal.""" """Test filename validation to prevent path traversal."""
mock_completion = MagicMock() mock_provider = MagicMock()
# Try to inject a malicious filename # 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"} {"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" 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.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @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")
def test_validates_filename_with_dots(self, mock_client, mock_log_progress, mock_embed_task): def test_validates_filename_with_dots(self, mock_get_provider, mock_log_progress, mock_embed_task):
"""Test filename validation rejects '..' in filenames.""" """Test filename validation rejects '..' in filenames."""
mock_completion = MagicMock() mock_provider = MagicMock()
mock_completion.choices[0].message.content = json.dumps( mock_provider.chat_completion.return_value = json.dumps(
{"filename": "test..invoice.pdf", "document_type": "Invoice"} {"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" 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.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @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")
def test_accepts_valid_filename(self, mock_client, mock_log_progress, mock_embed_task): def test_accepts_valid_filename(self, mock_get_provider, mock_log_progress, mock_embed_task):
"""Test that valid filenames are accepted.""" """Test that valid filenames are accepted."""
mock_completion = MagicMock() mock_provider = MagicMock()
mock_completion.choices[0].message.content = json.dumps( mock_provider.chat_completion.return_value = json.dumps(
{"filename": "2024-01-15_Invoice_Amazon.pdf", "document_type": "Invoice"} {"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" 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.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @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")
def test_handles_malformed_json_with_valid_structure(self, mock_client, mock_log_progress, mock_embed_task): 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.""" """Test handling of JSON that's parseable but missing expected fields."""
mock_completion = MagicMock() mock_provider = MagicMock()
mock_completion.choices[0].message.content = '{"unexpected_field": "value"}' mock_provider.chat_completion.return_value = '{"unexpected_field": "value"}'
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" 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.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @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")
def test_handles_absolute_path_filename(self, mock_client, mock_log_progress, mock_embed_task): 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 (line 73).""" """Test handling when filename is provided as an absolute path."""
mock_completion = MagicMock() mock_provider = MagicMock()
mock_completion.choices[0].message.content = '{"filename": "test.pdf", "document_type": "Unknown"}' mock_provider.chat_completion.return_value = '{"filename": "test.pdf", "document_type": "Unknown"}'
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" 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.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @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") @patch("app.tasks.extract_metadata_with_gpt.SessionLocal")
def test_database_lookup_with_existing_file( 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).""" """Test file_id retrieval when file exists on disk and in database."""
mock_completion = MagicMock() mock_provider = MagicMock()
mock_completion.choices[0].message.content = '{"filename": "test.pdf", "document_type": "Unknown"}' mock_provider.chat_completion.return_value = '{"filename": "test.pdf", "document_type": "Unknown"}'
mock_client.chat.completions.create.return_value = mock_completion mock_get_provider.return_value = mock_provider
# Mock database session # Mock database session
mock_db = MagicMock() mock_db = MagicMock()
@@ -333,20 +334,14 @@ class TestExtractMetadataWithGpt:
@pytest.mark.unit @pytest.mark.unit
class TestClientInitialization: class TestModuleImports:
"""Tests for OpenAI client initialization error handling.""" """Tests for module import behaviour."""
def test_client_initialization_imports_successfully(self): def test_module_imports_successfully(self):
"""Test that module imports successfully even if client initialization fails (lines 25-27). """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 assert mod is not None
on failure. This test verifies the module can be imported without crashing, assert hasattr(mod, "extract_metadata_with_gpt")
regardless of whether the client initializes successfully or not. assert hasattr(mod, "extract_json_from_text")
""" assert hasattr(mod, "get_ai_provider")
# 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
+18 -23
View File
@@ -407,65 +407,60 @@ startxref
@pytest.mark.unit @pytest.mark.unit
class TestRefineTextWithGPT: 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") @patch("app.tasks.refine_text_with_gpt.log_task_progress")
def test_successful_text_refinement(self, mock_log): 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" raw_text = "This is s0me text with OCR err0rs"
filename = "test.pdf" filename = "test.pdf"
cleaned = "This is some text with OCR errors"
# 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
# Import the module to patch the correct function # Import the module to patch the correct function
from app.tasks import extract_metadata_with_gpt as metadata_module from app.tasks import extract_metadata_with_gpt as metadata_module
mock_provider = MagicMock()
mock_provider.chat_completion.return_value = cleaned
with ( 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.object(metadata_module, "extract_metadata_with_gpt") as mock_extract,
patch("app.tasks.refine_text_with_gpt.settings") as mock_settings, patch("app.tasks.refine_text_with_gpt.settings") as mock_settings,
): ):
mock_settings.openai_model = "gpt-4" mock_settings.openai_model = "gpt-4"
mock_settings.ai_model = None
mock_extract.delay = MagicMock() mock_extract.delay = MagicMock()
result = refine_text_with_gpt.run(filename, raw_text) result = refine_text_with_gpt.run(filename, raw_text)
# Verify results # Verify results
assert result["filename"] == filename 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 # Verify AI provider was called correctly
mock_client.chat.completions.create.assert_called_once() mock_provider.chat_completion.assert_called_once()
call_kwargs = mock_client.chat.completions.create.call_args[1] call_kwargs = mock_provider.chat_completion.call_args[1]
assert call_kwargs["model"] == "gpt-4"
assert len(call_kwargs["messages"]) == 2 assert len(call_kwargs["messages"]) == 2
assert call_kwargs["messages"][1]["content"] == raw_text assert call_kwargs["messages"][1]["content"] == raw_text
# Verify metadata extraction was queued # 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") @patch("app.tasks.refine_text_with_gpt.log_task_progress")
def test_openai_api_error(self, mock_log): 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" raw_text = "Test text"
filename = "test.pdf" filename = "test.pdf"
mock_client = Mock() mock_provider = MagicMock()
mock_client.chat.completions.create.side_effect = Exception("OpenAI API error") mock_provider.chat_completion.side_effect = Exception("OpenAI API error")
with ( 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, patch("app.tasks.refine_text_with_gpt.settings") as mock_settings,
): ):
mock_settings.openai_model = "gpt-4" mock_settings.openai_model = "gpt-4"
mock_settings.ai_model = None
# Should raise the exception # Should raise the exception
with pytest.raises(Exception) as exc_info: with pytest.raises(Exception) as exc_info: