diff --git a/app/config.py b/app/config.py index 0def557c..f9a18df5 100644 --- a/app/config.py +++ b/app/config.py @@ -15,6 +15,28 @@ class Settings(BaseSettings): openai_api_key: str openai_base_url: str = "https://api.openai.com/v1" # Default to OpenAI's endpoint 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 debug: bool = False # Default to False diff --git a/app/tasks/extract_metadata_with_gpt.py b/app/tasks/extract_metadata_with_gpt.py index c6fd33f9..162a75f3 100644 --- a/app/tasks/extract_metadata_with_gpt.py +++ b/app/tasks/extract_metadata_with_gpt.py @@ -5,8 +5,6 @@ import logging import os import re -import openai - # Import the shared Celery instance from app.celery_app import celery 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.retry_config import BaseTaskWithRetry from app.utils import log_task_progress +from app.utils.ai_provider import get_ai_provider 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): """ @@ -114,23 +105,24 @@ def extract_metadata_with_gpt(self, filename: str, cleaned_text: str, file_id: i try: 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) - completion = client.chat.completions.create( - model=settings.openai_model, + log_task_progress(task_id, "call_ai_provider", "in_progress", "Calling AI provider API", file_id=file_id) + provider = get_ai_provider() + model = settings.ai_model or settings.openai_model + content = provider.chat_completion( messages=[ {"role": "system", "content": "You are an intelligent document classifier."}, {"role": "user", "content": prompt}, ], + model=model, temperature=0, ) - content = completion.choices[0].message.content logger.info(f"[{task_id}] Raw classification response for {filename}: {content[:200]}...") log_task_progress( task_id, - "call_openai", + "call_ai_provider", "success", - "Received OpenAI response", + "Received AI provider response", file_id=file_id, 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} 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( task_id, "extract_metadata_with_gpt", "failure", f"Exception: {str(e)}", 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 {} diff --git a/app/tasks/refine_text_with_gpt.py b/app/tasks/refine_text_with_gpt.py index d67f4c94..62b74c14 100644 --- a/app/tasks/refine_text_with_gpt.py +++ b/app/tasks/refine_text_with_gpt.py @@ -2,32 +2,29 @@ import logging -import openai - # Import the shared Celery instance from app.celery_app import celery from app.config import settings from app.tasks.retry_config import BaseTaskWithRetry from app.utils import log_task_progress +from app.utils.ai_provider import get_ai_provider 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) 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 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}") 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( - model=settings.openai_model, + provider = get_ai_provider() + model = settings.ai_model or settings.openai_model + cleaned_text = provider.chat_completion( messages=[ { "role": "system", @@ -38,16 +35,15 @@ def refine_text_with_gpt(self, filename: str, raw_text: str): }, {"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") log_task_progress( task_id, - "call_openai", + "call_ai_provider", "success", - "Received refined text from OpenAI", + "Received refined text from AI provider", detail=f"Input: {len(raw_text)} chars → Output: {len(cleaned_text)} chars", ) diff --git a/app/utils/ai_provider.py b/app/utils/ai_provider.py new file mode 100644 index 00000000..a98cae90 --- /dev/null +++ b/app/utils/ai_provider.py @@ -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 (0–1). 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" + ) diff --git a/requirements.txt b/requirements.txt index b1389fd6..a46fff38 100644 --- a/requirements.txt +++ b/requirements.txt @@ -34,4 +34,7 @@ boto3>=1.28.0 paramiko>=3.4.0 # SSH/SFTP implementation for Python (LGPL license) # 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 diff --git a/tests/test_ai_provider.py b/tests/test_ai_provider.py new file mode 100644 index 00000000..e77caefa --- /dev/null +++ b/tests/test_ai_provider.py @@ -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 diff --git a/tests/test_extract_metadata_gpt.py b/tests/test_extract_metadata_gpt.py index a39325e0..ff96bad8 100644 --- a/tests/test_extract_metadata_gpt.py +++ b/tests/test_extract_metadata_gpt.py @@ -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") diff --git a/tests/test_ocr_processing.py b/tests/test_ocr_processing.py index 29d3af20..c2c9159e 100644 --- a/tests/test_ocr_processing.py +++ b/tests/test_ocr_processing.py @@ -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: