From 0736cd8710ba4864e745ebafadcf165b0693af43 Mon Sep 17 00:00:00 2001
From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com>
Date: Mon, 23 Feb 2026 19:17:09 +0000
Subject: [PATCH 1/5] Initial plan
From d4c7fb26ac6a8365843310979ee83f7e1538cf95 Mon Sep 17 00:00:00 2001
From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com>
Date: Mon, 23 Feb 2026 19:26:45 +0000
Subject: [PATCH 2/5] 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>
---
app/config.py | 22 +
app/tasks/extract_metadata_with_gpt.py | 28 +-
app/tasks/refine_text_with_gpt.py | 22 +-
app/utils/ai_provider.py | 345 +++++++++++++++
requirements.txt | 5 +-
tests/test_ai_provider.py | 591 +++++++++++++++++++++++++
tests/test_extract_metadata_gpt.py | 157 ++++---
tests/test_ocr_processing.py | 41 +-
8 files changed, 1075 insertions(+), 136 deletions(-)
create mode 100644 app/utils/ai_provider.py
create mode 100644 tests/test_ai_provider.py
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:
From dcdef443032e5c43202566b3545560c3464f3848 Mon Sep 17 00:00:00 2001
From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com>
Date: Mon, 23 Feb 2026 19:42:03 +0000
Subject: [PATCH 3/5] feat: add Portkey provider support and null-content guard
to AI abstraction layer
Co-authored-by: christianlouis <361235+christianlouis@users.noreply.github.com>
---
app/config.py | 7 ++
app/utils/ai_provider.py | 131 ++++++++++++++++++++--
docs/ConfigurationGuide.md | 147 +++++++++++++++++++++++-
tests/test_ai_provider.py | 173 +++++++++++++++++++++++++++++
tests/test_extract_metadata_gpt.py | 22 ++--
tests/test_ocr_processing.py | 4 +-
6 files changed, 457 insertions(+), 27 deletions(-)
diff --git a/app/config.py b/app/config.py
index f9a18df5..702d6485 100644
--- a/app/config.py
+++ b/app/config.py
@@ -35,6 +35,13 @@ class Settings(BaseSettings):
openrouter_api_key: Optional[str] = None
openrouter_base_url: str = "https://openrouter.ai/api/v1"
+ # Portkey AI gateway settings (used when ai_provider="portkey")
+ # See https://portkey.ai for setup instructions
+ portkey_api_key: Optional[str] = None
+ portkey_virtual_key: Optional[str] = None # Routes to a specific provider via Portkey vault
+ portkey_config: Optional[str] = None # Portkey Config ID for advanced routing rules
+ portkey_base_url: str = "https://api.portkey.ai/v1"
+
# Azure OpenAI API version (used when ai_provider="azure")
azure_openai_api_version: str = "2024-02-01"
workdir: str
diff --git a/app/utils/ai_provider.py b/app/utils/ai_provider.py
index a98cae90..fa9353b0 100644
--- a/app/utils/ai_provider.py
+++ b/app/utils/ai_provider.py
@@ -3,8 +3,8 @@
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.
+Google Gemini, Ollama (local LLMs), OpenRouter, Portkey, 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.
@@ -14,9 +14,36 @@ import logging
from abc import ABC, abstractmethod
from typing import Any, Dict, List, Optional
+from app.config import settings
+
logger = logging.getLogger(__name__)
+def _require_text_content(content: Optional[str]) -> str:
+ """Raise a clear error if the AI response contains no text content.
+
+ This can happen when the model returns a tool/function call instead of
+ a plain text message. All DocuElevate prompts expect a plain-text or
+ JSON response, so ``None`` content is always an unexpected condition.
+
+ Args:
+ content: The ``message.content`` value from the completion response.
+
+ Returns:
+ The original string, guaranteed non-None.
+
+ Raises:
+ ValueError: If *content* is ``None``.
+ """
+ if content is None:
+ raise ValueError(
+ "AI provider returned a response with no text content (content=None). "
+ "This may occur when the model generates a tool call instead of a plain text reply. "
+ "Ensure the model and prompt are configured for text/JSON output."
+ )
+ return content
+
+
class AIProvider(ABC):
"""Abstract base class for AI chat completion providers.
@@ -80,7 +107,8 @@ class OpenAIProvider(AIProvider):
temperature=temperature,
**kwargs,
)
- return completion.choices[0].message.content
+ _content = completion.choices[0].message.content
+ return _require_text_content(_content)
class AzureOpenAIProvider(AIProvider):
@@ -108,7 +136,8 @@ class AzureOpenAIProvider(AIProvider):
temperature=temperature,
**kwargs,
)
- return completion.choices[0].message.content
+ _content = completion.choices[0].message.content
+ return _require_text_content(_content)
class AnthropicProvider(AIProvider):
@@ -139,7 +168,8 @@ class AnthropicProvider(AIProvider):
api_key=self._api_key,
**kwargs,
)
- return response.choices[0].message.content
+ _content = response.choices[0].message.content
+ return _require_text_content(_content)
class GeminiProvider(AIProvider):
@@ -170,7 +200,8 @@ class GeminiProvider(AIProvider):
api_key=self._api_key,
**kwargs,
)
- return response.choices[0].message.content
+ _content = response.choices[0].message.content
+ return _require_text_content(_content)
class OllamaProvider(AIProvider):
@@ -210,7 +241,8 @@ class OllamaProvider(AIProvider):
temperature=temperature,
**kwargs,
)
- return completion.choices[0].message.content
+ _content = completion.choices[0].message.content
+ return _require_text_content(_content)
class OpenRouterProvider(AIProvider):
@@ -243,7 +275,74 @@ class OpenRouterProvider(AIProvider):
temperature=temperature,
**kwargs,
)
- return completion.choices[0].message.content
+ _content = completion.choices[0].message.content
+ return _require_text_content(_content)
+
+
+class PortkeyProvider(AIProvider):
+ """Portkey AI gateway (https://portkey.ai).
+
+ Portkey is an AI gateway that provides observability, caching, automatic
+ retries, fallbacks, and load balancing across 200+ LLMs via a single
+ OpenAI-compatible endpoint.
+
+ Required settings:
+ ``PORTKEY_API_KEY`` – your Portkey account API key.
+
+ Optional settings:
+ ``PORTKEY_VIRTUAL_KEY`` – a Portkey *Virtual Key* that maps to the
+ credentials of a specific provider stored in your Portkey vault.
+ When set, you do not need to expose the underlying provider's API key
+ in your environment.
+
+ ``PORTKEY_CONFIG`` – a saved Portkey *Config* ID (e.g.
+ ``pc-my-config-abc123``) that applies advanced routing rules such as
+ fallbacks and load balancing.
+
+ ``PORTKEY_BASE_URL`` – override the gateway endpoint.
+ Default: ``https://api.portkey.ai/v1``.
+
+ The model name should match what the underlying provider expects (e.g.
+ ``gpt-4o`` for OpenAI, ``claude-3-5-sonnet-20241022`` for Anthropic via a
+ virtual key).
+ """
+
+ def __init__(
+ self,
+ api_key: str,
+ virtual_key: Optional[str] = None,
+ config: Optional[str] = None,
+ base_url: str = "https://api.portkey.ai/v1",
+ ) -> None:
+ import openai
+
+ portkey_headers: Dict[str, str] = {"x-portkey-api-key": api_key}
+ if virtual_key:
+ portkey_headers["x-portkey-virtual-key"] = virtual_key
+ if config:
+ portkey_headers["x-portkey-config"] = config
+
+ self._client = openai.OpenAI(
+ api_key=api_key,
+ base_url=base_url,
+ default_headers=portkey_headers,
+ )
+
+ 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,
+ )
+ _content = completion.choices[0].message.content
+ return _require_text_content(_content)
class LiteLLMProvider(AIProvider):
@@ -283,7 +382,8 @@ class LiteLLMProvider(AIProvider):
completion_kwargs["api_base"] = self._api_base
completion_kwargs.update(kwargs)
response = litellm.completion(**completion_kwargs)
- return response.choices[0].message.content
+ _content = response.choices[0].message.content
+ return _require_text_content(_content)
def get_ai_provider() -> AIProvider:
@@ -300,8 +400,6 @@ def get_ai_provider() -> AIProvider:
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}")
@@ -333,6 +431,15 @@ def get_ai_provider() -> AIProvider:
api_key=settings.openrouter_api_key,
base_url=settings.openrouter_base_url,
)
+ elif provider == "portkey":
+ if not settings.portkey_api_key:
+ raise ValueError("PORTKEY_API_KEY must be set when AI_PROVIDER='portkey'")
+ return PortkeyProvider(
+ api_key=settings.portkey_api_key,
+ virtual_key=settings.portkey_virtual_key,
+ config=settings.portkey_config,
+ base_url=settings.portkey_base_url,
+ )
elif provider == "litellm":
return LiteLLMProvider(
api_key=settings.openai_api_key or None,
@@ -341,5 +448,5 @@ def get_ai_provider() -> AIProvider:
else:
raise ValueError(
f"Unknown AI provider: '{provider}'. "
- "Supported providers: openai, azure, anthropic, gemini, ollama, openrouter, litellm"
+ "Supported providers: openai, azure, anthropic, gemini, ollama, openrouter, portkey, litellm"
)
diff --git a/docs/ConfigurationGuide.md b/docs/ConfigurationGuide.md
index de57935b..b44cbbae 100644
--- a/docs/ConfigurationGuide.md
+++ b/docs/ConfigurationGuide.md
@@ -302,11 +302,154 @@ SECURITY_HEADER_CSP_VALUE="default-src 'self'; script-src 'self' https://trusted
- [Deployment Guide - Security Headers](DeploymentGuide.md#security-headers) for Traefik/Nginx examples
- [SECURITY_AUDIT.md](../SECURITY_AUDIT.md#infrastructure-security) for security rationale
-### OpenAI & Azure Document Intelligence
+### AI Provider & Model Selection
+
+DocuElevate supports multiple AI providers for metadata extraction and OCR text refinement. Select the provider via `AI_PROVIDER` and configure the matching credentials below.
+
+| **Variable** | **Description** | **Default** |
+|-------------------|-----------------------------------------------------------------------|--------------------|
+| `AI_PROVIDER` | Active AI provider. See supported values below. | `openai` |
+| `AI_MODEL` | Model name for the selected provider. Falls back to `OPENAI_MODEL` when not set. | *(unset)* |
+| `OPENAI_MODEL` | Default model name (used when `AI_MODEL` is not set). | `gpt-4o-mini` |
+
+**Supported `AI_PROVIDER` values**: `openai`, `azure`, `anthropic`, `gemini`, `ollama`, `openrouter`, `portkey`, `litellm`
+
+---
+
+#### OpenAI (default)
+
+| **Variable** | **Description** | **Default** |
+|-----------------------|--------------------------------------------------|----------------------------------|
+| `OPENAI_API_KEY` | OpenAI API key. | *(required)* |
+| `OPENAI_BASE_URL` | API base URL. Change for compatible proxies. | `https://api.openai.com/v1` |
+
+```bash
+AI_PROVIDER=openai
+OPENAI_API_KEY=sk-...
+OPENAI_MODEL=gpt-4o-mini
+```
+
+#### Azure OpenAI
+
+| **Variable** | **Description** | **Default** |
+|-------------------------------|----------------------------------------------|----------------|
+| `OPENAI_API_KEY` | Azure OpenAI API key. | *(required)* |
+| `OPENAI_BASE_URL` | Azure resource endpoint URL. | *(required)* |
+| `AZURE_OPENAI_API_VERSION` | Azure OpenAI API version string. | `2024-02-01` |
+
+```bash
+AI_PROVIDER=azure
+OPENAI_API_KEY=
- We harness the power of OpenAI for metadata extraction and text refinement, integrate seamlessly + We harness the power of pluggable AI providers (OpenAI, Anthropic Claude, Google Gemini, Ollama, + OpenRouter, Portkey, and more) for metadata extraction and text refinement, integrate seamlessly with Dropbox, Nextcloud, and Paperless NGX for storage and indexing, leverage Azure Document Intelligence for OCR, and even use Gotenberg for file-to-PDF conversions.
@@ -33,7 +34,7 @@