Merge pull request #374 from christianlouis/copilot/build-ai-provider-abstraction

feat: pluggable AI provider abstraction layer (OpenAI, Azure, Anthropic, Gemini, Ollama, OpenRouter, Portkey, LiteLLM)
This commit is contained in:
Christian Krakau-Louis
2026-02-23 21:35:35 +01:00
committed by GitHub
26 changed files with 1900 additions and 251 deletions
+33 -5
View File
@@ -119,16 +119,44 @@ AUTHENTIK_CONFIG_URL=<ConfigUrlOfYourApp, e.g. https://authentik.example.com/app
OAUTH_PROVIDER_NAME="Authentik SSO"
# **AI/ML Services**
# OpenAI
# Select your AI provider: openai | azure | anthropic | gemini | ollama | openrouter | portkey | litellm
AI_PROVIDER=openai
# Model override (optional falls back to OPENAI_MODEL when not set)
# AI_MODEL=gpt-4o-mini
# --- OpenAI (AI_PROVIDER=openai) ---
OPENAI_API_KEY="<OPENAI_API_KEY>"
OPENAI_BASE_URL=https://api.openai.com/v1
OPENAI_MODEL=gpt-4o-mini
# Azure AI
AZURE_REGION="eastus"
AZURE_ENDPOINT="https://<yourendpoint>.cognitiveservices.azure.com/"
AZURE_AI_KEY=<AZURE_AI_KEY>
# --- Anthropic Claude (AI_PROVIDER=anthropic) ---
# ANTHROPIC_API_KEY=sk-ant-...
# AI_MODEL=claude-3-5-sonnet-20241022
# --- Google Gemini (AI_PROVIDER=gemini) ---
# GEMINI_API_KEY=AIza...
# AI_MODEL=gemini-1.5-pro
# --- Ollama local LLMs (AI_PROVIDER=ollama) ---
# OLLAMA_BASE_URL=http://localhost:11434
# AI_MODEL=llama3.2
# --- OpenRouter (AI_PROVIDER=openrouter) ---
# OPENROUTER_API_KEY=sk-or-...
# AI_MODEL=anthropic/claude-3.5-sonnet
# --- Portkey AI Gateway (AI_PROVIDER=portkey) ---
# PORTKEY_API_KEY=pk-...
# PORTKEY_VIRTUAL_KEY=vk-... # optional routes to provider credentials in Portkey vault
# PORTKEY_CONFIG=pc-... # optional saved Config ID for fallbacks / load balancing
# --- Azure OpenAI (AI_PROVIDER=azure) ---
# OPENAI_API_KEY=<azure-key>
# OPENAI_BASE_URL=https://my-resource.openai.azure.com
# AZURE_OPENAI_API_VERSION=2024-02-01
# AI_MODEL=gpt-4o # deployment name in Azure
# Azure Document Intelligence (OCR separate from AI provider above)
# **Email Settings**
EMAIL_HOST=smtp.example.com
EMAIL_PORT=587
+6 -5
View File
@@ -31,7 +31,7 @@
DocuElevate automates the handling, extraction, and processing of documents using a variety of services, including:
- **OpenAI** for metadata extraction and text refinement.
- **AI Provider** (pluggable OpenAI, Anthropic, Gemini, Ollama, OpenRouter, Portkey, and more) for metadata extraction and text refinement.
- **Dropbox**, **Nextcloud**, and **Google Drive** for file storage and uploads.
- **Paperless NGX** for document indexing and management.
- **Azure Document Intelligence** for OCR on PDFs.
@@ -86,7 +86,7 @@ Documents enter DocuElevate through four possible channels:
Every document goes through the following steps:
1. **PDF Conversion**: Non-PDF files are converted to PDF format using Gotenberg
2. **OCR Processing**: Azure Document Intelligence extracts text from images/scans
3. **Metadata Extraction**: OpenAI analyzes document content to identify:
3. **Metadata Extraction**: The configured AI provider analyzes document content to identify:
- Document type (invoice, receipt, contract, etc.)
- Key entities (dates, names, amounts, account numbers)
- Important data points specific to the document type
@@ -116,8 +116,8 @@ Users can choose to send documents to any combination of these destinations thro
- Manual uploads (via API or UI) to Dropbox, Nextcloud, Google Drive, or Paperless
- **OCR Processing (Azure)**:
- Extract text from scanned PDFs using Azure Document Intelligence
- **Metadata Extraction (OpenAI)**:
- Use GPT to classify, label, or otherwise enrich the text with structured metadata
- **Metadata Extraction (AI Provider)**:
- Use any supported AI provider (OpenAI, Anthropic, Gemini, Ollama, etc.) to classify, label, or otherwise enrich the text with structured metadata
- **PDF Conversion (Gotenberg)**:
- Convert non-PDF attachments (e.g., Word docs, images) into PDFs
- **Document Management (Paperless NGX)**:
@@ -214,7 +214,8 @@ The following is a summary of the licenses used by our direct dependencies:
| Uvicorn | BSD |
| SQLAlchemy | MIT |
| Pydantic | MIT |
| OpenAI | MIT |
| openai | MIT |
| litellm | MIT |
| pypdf | BSD |
| Requests | Apache 2.0 |
| puremagic | MIT |
+57 -1
View File
@@ -1,5 +1,9 @@
"""
OpenAI API endpoints
AI provider and OpenAI API endpoints.
Exposes two endpoints:
- GET /api/ai/test tests the currently configured AI provider (generic, provider-agnostic)
- GET /api/openai/test backward-compatible alias that tests the OpenAI API specifically
"""
import logging
@@ -146,3 +150,55 @@ async def test_openai_connection(request: Request):
except Exception as e:
logger.exception("Unexpected error testing OpenAI connection")
return {"status": "error", "message": f"Unexpected error: {str(e)}"}
@router.get("/ai/test")
@require_login
async def test_ai_provider_connection(request: Request):
"""
Test the currently configured AI provider connection.
Uses ``get_ai_provider()`` to instantiate the active provider and sends a
minimal chat completion to verify that the credentials and endpoint are
reachable. Works for all supported providers (OpenAI, Azure, Anthropic,
Gemini, Ollama, OpenRouter, Portkey, LiteLLM).
"""
from app.utils.ai_provider import get_ai_provider
provider_name = settings.ai_provider
model = settings.ai_model or settings.openai_model
logger.info(f"Testing AI provider connection: provider={provider_name}, model={model}")
try:
provider = get_ai_provider()
response = provider.chat_completion(
messages=[{"role": "user", "content": "Reply with the single word: ok"}],
model=model,
temperature=0,
max_tokens=5,
)
logger.info(f"AI provider test successful: provider={provider_name}")
return {
"status": "success",
"message": f"AI provider '{provider_name}' is reachable and responding",
"provider": provider_name,
"model": model,
"response_preview": (response or "")[:50],
}
except ValueError as e:
# Configuration errors (missing keys, unknown provider)
logger.warning(f"AI provider configuration error: {e}")
return {
"status": "error",
"message": str(e),
"provider": provider_name,
}
except Exception as e:
detail = _get_exception_chain_detail(e)
logger.error(f"AI provider test failed for '{provider_name}': {detail}", exc_info=True)
return {
"status": "error",
"message": f"Connection failed: {detail}",
"provider": provider_name,
}
+29
View File
@@ -15,6 +15,35 @@ 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"
# 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
debug: bool = False # Default to False
+15 -6
View File
@@ -11,7 +11,7 @@ from app.api.google_drive import test_google_drive_token
from app.api.onedrive import test_onedrive_token
# Import the test functions from API routes
from app.api.openai import test_openai_connection
from app.api.openai import test_ai_provider_connection, test_openai_connection
from app.celery_app import celery
from app.config import settings
@@ -75,8 +75,17 @@ def unwrap_decorated_function(func):
# Create synchronous versions of the test functions that bypass authentication
def sync_test_ai_provider_connection():
"""Synchronous wrapper for the AI provider test function that bypasses auth."""
inner_func = unwrap_decorated_function(test_ai_provider_connection)
request = MockRequest()
if inspect.iscoroutinefunction(inner_func):
return asyncio.run(inner_func(request))
return inner_func(request)
def sync_test_openai_connection():
"""Synchronous wrapper for the OpenAI test function that bypasses auth"""
"""Synchronous wrapper for the OpenAI test function that bypasses auth."""
# Get the original function without the @require_login decorator
inner_func = unwrap_decorated_function(test_openai_connection)
request = MockRequest()
@@ -139,10 +148,10 @@ def check_credentials():
# Define services with their test functions and configuration status
services = [
{
"name": "OpenAI",
"check_func": sync_test_openai_connection,
"configured": provider_status.get("OpenAI", {}).get("configured", False),
"config_issues": [], # OpenAI isn't in storage_configs
"name": "AI Provider",
"check_func": sync_test_ai_provider_connection,
"configured": provider_status.get("AI Provider", {}).get("configured", False),
"config_issues": [],
},
{
"name": "Azure Document Intelligence",
+10 -18
View File
@@ -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 {}
+9 -13
View File
@@ -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",
)
+452
View File
@@ -0,0 +1,452 @@
#!/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, 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.
"""
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.
All concrete providers must implement :meth:`chat_completion`, which
accepts a list of chat messages and returns the model's response as a
plain string. The interface intentionally mirrors the OpenAI Chat
Completions API so that callers need no provider-specific knowledge.
"""
@abstractmethod
def chat_completion(
self,
messages: List[Dict[str, str]],
model: str,
temperature: float = 0,
**kwargs: Any,
) -> str:
"""Get a chat completion from the AI provider.
Args:
messages: List of message dicts with ``role`` and ``content`` keys.
model: Model name/identifier to use (provider-specific format).
temperature: Sampling temperature (01). Default: 0 (deterministic).
**kwargs: Additional provider-specific arguments passed through.
Returns:
The model's response as a plain string.
Raises:
Exception: If the underlying API call fails.
"""
class OpenAIProvider(AIProvider):
"""OpenAI provider using the ``openai`` Python SDK.
Also works as a drop-in for any OpenAI-compatible API endpoint, including
LocalAI and LM Studio. Ollama and OpenRouter have dedicated providers with
sensible defaults, but this provider works for them too when a custom
``base_url`` is supplied.
"""
def __init__(self, api_key: str, base_url: Optional[str] = None) -> None:
import openai
self._client = openai.OpenAI(
api_key=api_key,
base_url=base_url or "https://api.openai.com/v1",
)
def chat_completion(
self,
messages: List[Dict[str, str]],
model: str,
temperature: float = 0,
**kwargs: Any,
) -> str:
completion = self._client.chat.completions.create(
model=model,
messages=messages,
temperature=temperature,
**kwargs,
)
_content = completion.choices[0].message.content
return _require_text_content(_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,
)
_content = completion.choices[0].message.content
return _require_text_content(_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,
)
_content = response.choices[0].message.content
return _require_text_content(_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,
)
_content = response.choices[0].message.content
return _require_text_content(_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,
)
_content = completion.choices[0].message.content
return _require_text_content(_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,
)
_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):
"""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)
_content = response.choices[0].message.content
return _require_text_content(_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.
"""
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 == "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,
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, portkey, litellm"
)
+62 -13
View File
@@ -57,20 +57,69 @@ def get_provider_status() -> dict[str, dict[str, object]]:
"test_endpoint": "/api/diagnostic/test-notification",
}
# Add AI services first
providers["OpenAI"] = {
"name": "OpenAI",
"icon": "fa-brands fa-openai",
"configured": bool(
getattr(settings, "openai_api_key", None) and str(getattr(settings, "openai_api_key", "")).startswith("sk-")
),
# Add AI provider (dynamic reflects whichever provider is configured)
ai_provider_name = getattr(settings, "ai_provider", "openai").lower()
model = getattr(settings, "ai_model", None) or getattr(settings, "openai_model", "gpt-4o-mini")
# Determine whether the active provider has its required credentials set
def _ai_configured() -> bool:
if ai_provider_name in ("openai", "azure", "litellm"):
return bool(getattr(settings, "openai_api_key", None))
elif ai_provider_name == "anthropic":
return bool(getattr(settings, "anthropic_api_key", None))
elif ai_provider_name == "gemini":
return bool(getattr(settings, "gemini_api_key", None))
elif ai_provider_name == "ollama":
return bool(getattr(settings, "ollama_base_url", None))
elif ai_provider_name == "openrouter":
return bool(getattr(settings, "openrouter_api_key", None))
elif ai_provider_name == "portkey":
return bool(getattr(settings, "portkey_api_key", None))
return False
# Build provider-specific detail rows
def _ai_details() -> dict:
base = {"provider": ai_provider_name, "model": model}
if ai_provider_name in ("openai", "azure", "litellm"):
base["api_key"] = mask_sensitive_value(getattr(settings, "openai_api_key", None))
base["base_url"] = getattr(settings, "openai_base_url", "https://api.openai.com/v1")
elif ai_provider_name == "anthropic":
base["api_key"] = mask_sensitive_value(getattr(settings, "anthropic_api_key", None))
elif ai_provider_name == "gemini":
base["api_key"] = mask_sensitive_value(getattr(settings, "gemini_api_key", None))
elif ai_provider_name == "ollama":
base["base_url"] = getattr(settings, "ollama_base_url", "http://localhost:11434")
elif ai_provider_name == "openrouter":
base["api_key"] = mask_sensitive_value(getattr(settings, "openrouter_api_key", None))
base["base_url"] = getattr(settings, "openrouter_base_url", "https://openrouter.ai/api/v1")
elif ai_provider_name == "portkey":
base["api_key"] = mask_sensitive_value(getattr(settings, "portkey_api_key", None))
if getattr(settings, "portkey_virtual_key", None):
base["virtual_key"] = mask_sensitive_value(settings.portkey_virtual_key)
if getattr(settings, "portkey_config", None):
base["config"] = settings.portkey_config
return base
_PROVIDER_ICONS = {
"openai": "fa-brands fa-openai",
"azure": "fa-brands fa-microsoft",
"anthropic": "fa-solid fa-robot",
"gemini": "fa-brands fa-google",
"ollama": "fa-solid fa-server",
"openrouter": "fa-solid fa-route",
"portkey": "fa-solid fa-key",
"litellm": "fa-solid fa-layer-group",
}
providers["AI Provider"] = {
"name": "AI Provider",
"icon": _PROVIDER_ICONS.get(ai_provider_name, "fa-solid fa-microchip"),
"configured": _ai_configured(),
"enabled": True,
"description": "AI-powered document analysis and metadata extraction",
"details": {
"api_key": mask_sensitive_value(getattr(settings, "openai_api_key", None)),
"base_url": getattr(settings, "openai_base_url", "https://api.openai.com/v1"),
"model": getattr(settings, "openai_model", "gpt-4"),
},
"description": f"AI-powered metadata extraction and OCR refinement ({ai_provider_name})",
"details": _ai_details(),
"testable": True,
"test_endpoint": "/api/ai/test",
}
providers["Azure AI"] = {
@@ -170,9 +170,21 @@ def get_settings_for_display(show_values: bool = False) -> dict[str, list[dict[s
"s3_acl",
],
"AI Services": [
"ai_provider",
"ai_model",
"openai_api_key",
"openai_base_url",
"openai_model",
"anthropic_api_key",
"gemini_api_key",
"ollama_base_url",
"openrouter_api_key",
"openrouter_base_url",
"portkey_api_key",
"portkey_virtual_key",
"portkey_config",
"portkey_base_url",
"azure_openai_api_version",
"azure_ai_key",
"azure_endpoint",
"azure_region",
+109 -11
View File
@@ -136,15 +136,15 @@ SETTING_METADATA = {
# AI Services
"openai_api_key": {
"category": "AI Services",
"description": "OpenAI API key for metadata extraction",
"description": "API key for the AI provider (required for OpenAI, Azure, LiteLLM; unused for Ollama)",
"type": "string",
"sensitive": True,
"required": True,
"required": False,
"restart_required": False,
},
"openai_base_url": {
"category": "AI Services",
"description": "OpenAI API base URL (default: https://api.openai.com/v1)",
"description": "API base URL override for Azure endpoints, local proxies, or LiteLLM (default: https://api.openai.com/v1)",
"type": "string",
"sensitive": False,
"required": False,
@@ -152,34 +152,132 @@ SETTING_METADATA = {
},
"openai_model": {
"category": "AI Services",
"description": "OpenAI model to use (e.g., gpt-4o-mini)",
"description": "Fallback model name used when AI_MODEL is not set (e.g. gpt-4o-mini)",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": False,
},
"azure_ai_key": {
"ai_provider": {
"category": "AI Services",
"description": "Azure AI key for document intelligence",
"description": "Active AI provider for metadata extraction and OCR refinement",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": False,
"options": ["openai", "azure", "anthropic", "gemini", "ollama", "openrouter", "portkey", "litellm"],
},
"ai_model": {
"category": "AI Services",
"description": "Model name for the selected provider (overrides OPENAI_MODEL). E.g. gpt-4o, claude-3-5-sonnet-20241022, gemini-1.5-pro, llama3.2",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": False,
},
"anthropic_api_key": {
"category": "AI Services",
"description": "Anthropic API key (required when AI_PROVIDER=anthropic)",
"type": "string",
"sensitive": True,
"required": True,
"required": False,
"restart_required": False,
},
"gemini_api_key": {
"category": "AI Services",
"description": "Google AI Studio API key (required when AI_PROVIDER=gemini)",
"type": "string",
"sensitive": True,
"required": False,
"restart_required": False,
},
"ollama_base_url": {
"category": "AI Services",
"description": "Ollama server URL for local LLM inference (used when AI_PROVIDER=ollama)",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": False,
},
"openrouter_api_key": {
"category": "AI Services",
"description": "OpenRouter API key (required when AI_PROVIDER=openrouter)",
"type": "string",
"sensitive": True,
"required": False,
"restart_required": False,
},
"openrouter_base_url": {
"category": "AI Services",
"description": "OpenRouter gateway URL (default: https://openrouter.ai/api/v1)",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": False,
},
"portkey_api_key": {
"category": "AI Services",
"description": "Portkey account API key (required when AI_PROVIDER=portkey)",
"type": "string",
"sensitive": True,
"required": False,
"restart_required": False,
},
"portkey_virtual_key": {
"category": "AI Services",
"description": "Portkey Virtual Key routes to provider credentials stored in the Portkey vault (optional)",
"type": "string",
"sensitive": True,
"required": False,
"restart_required": False,
},
"portkey_config": {
"category": "AI Services",
"description": "Portkey Config ID for advanced routing, fallbacks, and load balancing (optional, e.g. pc-my-config-abc123)",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": False,
},
"portkey_base_url": {
"category": "AI Services",
"description": "Portkey gateway URL (default: https://api.portkey.ai/v1; override for self-hosted deployments)",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": False,
},
"azure_openai_api_version": {
"category": "AI Services",
"description": "Azure OpenAI API version string (used when AI_PROVIDER=azure, default: 2024-02-01)",
"type": "string",
"sensitive": False,
"required": False,
"restart_required": False,
},
# Azure Document Intelligence (OCR) separate from the AI provider above
"azure_ai_key": {
"category": "AI Services",
"description": "Azure Document Intelligence API key for OCR processing",
"type": "string",
"sensitive": True,
"required": False,
"restart_required": False,
},
"azure_region": {
"category": "AI Services",
"description": "Azure region for AI services",
"description": "Azure region for Document Intelligence services (e.g., eastus)",
"type": "string",
"sensitive": False,
"required": True,
"required": False,
"restart_required": False,
},
"azure_endpoint": {
"category": "AI Services",
"description": "Azure AI endpoint URL",
"description": "Azure Document Intelligence endpoint URL",
"type": "string",
"sensitive": False,
"required": True,
"required": False,
"restart_required": False,
},
# Storage Providers - Dropbox
+17 -28
View File
@@ -90,10 +90,21 @@ def get_required_settings() -> List[Dict[str, Any]]:
"wizard_step": 2,
"wizard_category": "Security",
},
{
"key": "ai_provider",
"label": "AI Provider",
"description": "AI provider for metadata extraction and OCR refinement (openai, azure, anthropic, gemini, ollama, openrouter, portkey, litellm)",
"type": "string",
"sensitive": False,
"default": "openai",
"options": ["openai", "azure", "anthropic", "gemini", "ollama", "openrouter", "portkey", "litellm"],
"wizard_step": 3,
"wizard_category": "AI Services",
},
{
"key": "openai_api_key",
"label": "OpenAI API Key",
"description": "API key for OpenAI services (metadata extraction)",
"label": "API Key (OpenAI / Azure / LiteLLM)",
"description": "API key for OpenAI, Azure OpenAI, or LiteLLM (not required for Ollama)",
"type": "string",
"sensitive": True,
"default": None,
@@ -101,32 +112,12 @@ def get_required_settings() -> List[Dict[str, Any]]:
"wizard_category": "AI Services",
},
{
"key": "azure_ai_key",
"label": "Azure AI Key",
"description": "Azure AI key for document intelligence (OCR)",
"type": "string",
"sensitive": True,
"default": None,
"wizard_step": 3,
"wizard_category": "AI Services",
},
{
"key": "azure_region",
"label": "Azure Region",
"description": "Azure region for AI services (e.g., eastus)",
"key": "openai_model",
"label": "Default Model",
"description": "Model name used when AI_MODEL is not set (e.g. gpt-4o-mini, claude-3-5-sonnet-20241022, llama3.2)",
"type": "string",
"sensitive": False,
"default": "eastus",
"wizard_step": 3,
"wizard_category": "AI Services",
},
{
"key": "azure_endpoint",
"label": "Azure Endpoint",
"description": "Azure AI endpoint URL",
"type": "string",
"sensitive": False,
"default": None,
"default": "gpt-4o-mini",
"wizard_step": 3,
"wizard_category": "AI Services",
},
@@ -147,8 +138,6 @@ def is_setup_required() -> bool:
critical_settings = [
("session_secret", ["INSECURE_DEFAULT_FOR_DEVELOPMENT_ONLY_DO_NOT_USE_IN_PRODUCTION_MINIMUM_32_CHARS"]),
("admin_password", [None, "", "your_secure_password", "changeme", "admin"]),
("openai_api_key", [None, "", "<OPENAI_API_KEY>", "test-key"]),
("azure_ai_key", [None, "", "<AZURE_AI_KEY>", "test-key"]),
]
for setting_key, invalid_values in critical_settings:
+145 -2
View File
@@ -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=<azure-key>
OPENAI_BASE_URL=https://my-resource.openai.azure.com
AI_MODEL=gpt-4o # deployment name in Azure
```
#### Anthropic Claude
| **Variable** | **Description** |
|---------------------|--------------------------|
| `ANTHROPIC_API_KEY` | Anthropic API key. |
```bash
AI_PROVIDER=anthropic
ANTHROPIC_API_KEY=sk-ant-...
AI_MODEL=claude-3-5-sonnet-20241022
```
#### Google Gemini
| **Variable** | **Description** |
|-------------------|----------------------------|
| `GEMINI_API_KEY` | Google AI Studio API key. |
```bash
AI_PROVIDER=gemini
GEMINI_API_KEY=AIza...
AI_MODEL=gemini-1.5-pro
```
#### Ollama (local LLMs CPU-friendly)
Run models locally using [Ollama](https://ollama.com). Recommended for CPU-only deployments:
| **Variable** | **Description** | **Default** |
|--------------------|-----------------------------------------|---------------------------|
| `OLLAMA_BASE_URL` | Ollama server URL. | `http://localhost:11434` |
```bash
AI_PROVIDER=ollama
OLLAMA_BASE_URL=http://ollama:11434 # Docker service name
AI_MODEL=llama3.2 # or qwen2.5, phi3, etc.
```
Recommended models for document processing on CPU:
- **`llama3.2`** (3B) good balance of speed and JSON output quality
- **`qwen2.5`** (3B/7B) excellent at structured extraction
- **`phi3`** (3.8B) strong reasoning, very fast on CPU
#### OpenRouter
[OpenRouter](https://openrouter.ai) provides access to 100+ models from a single endpoint using the `provider/model` name format.
| **Variable** | **Description** | **Default** |
|-------------------------|-------------------------------------|-----------------------------------|
| `OPENROUTER_API_KEY` | OpenRouter API key. | *(required)* |
| `OPENROUTER_BASE_URL` | Override the gateway URL. | `https://openrouter.ai/api/v1` |
```bash
AI_PROVIDER=openrouter
OPENROUTER_API_KEY=sk-or-...
AI_MODEL=anthropic/claude-3.5-sonnet
```
#### Portkey AI Gateway
[Portkey](https://portkey.ai) is an AI gateway that adds observability, caching, fallbacks, and load balancing across 200+ models behind a single OpenAI-compatible endpoint.
| **Variable** | **Description** | **Default** |
|-----------------------|----------------------------------------------------------------------------------------------------------|----------------------------------|
| `PORTKEY_API_KEY` | Portkey account API key. | *(required)* |
| `PORTKEY_VIRTUAL_KEY` | Optional Virtual Key (stores provider credentials in Portkey vault, keeping them out of your env file). | *(unset)* |
| `PORTKEY_CONFIG` | Optional saved Config ID (e.g. `pc-fallback-abc123`) for routing rules, fallbacks, and load balancing. | *(unset)* |
| `PORTKEY_BASE_URL` | Override the Portkey gateway URL (for self-hosted deployments). | `https://api.portkey.ai/v1` |
```bash
AI_PROVIDER=portkey
PORTKEY_API_KEY=pk-...
PORTKEY_VIRTUAL_KEY=vk-openai-abc123 # optional routes to your OpenAI key stored in Portkey
AI_MODEL=gpt-4o
```
Using a Config for fallback routing:
```bash
AI_PROVIDER=portkey
PORTKEY_API_KEY=pk-...
PORTKEY_CONFIG=pc-fallback-config-xyz # applies your saved routing rules
AI_MODEL=gpt-4o
```
#### LiteLLM (aggregator proxy)
[LiteLLM](https://litellm.ai) provides a unified `provider/model` interface for 100+ LLMs including OpenAI, Anthropic, Gemini, Cohere, Ollama, and many more.
| **Variable** | **Description** | **Default** |
|--------------------|-------------------------------------------------|-------------------------------|
| `OPENAI_API_KEY` | API key forwarded to LiteLLM (provider-specific). | *(depends on model)* |
| `OPENAI_BASE_URL` | Optional proxy/gateway URL. | `https://api.openai.com/v1` |
```bash
AI_PROVIDER=litellm
AI_MODEL=anthropic/claude-3-5-sonnet-20241022
OPENAI_API_KEY=sk-ant-... # passed as the api_key to LiteLLM
```
---
### Azure Document Intelligence
| **Variable** | **Description** | **How to Obtain** |
|---------------------------------|------------------------------------------|--------------------------------------------------------------------------|
| `OPENAI_API_KEY` | OpenAI API key for GPT metadata extraction. | [OpenAI API keys](https://platform.openai.com/account/api-keys) |
| `AZURE_DOCUMENT_INTELLIGENCE_KEY` | Azure Document Intelligence API key for OCR. | [Azure Portal](https://portal.azure.com/) |
| `AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT` | Endpoint URL for Azure Doc Intelligence API. | [Azure Portal](https://portal.azure.com/) |
+1 -1
View File
@@ -6,7 +6,7 @@ This guide provides instructions for deploying DocuElevate in various environmen
- Docker and Docker Compose
- Access to required external services (if configured):
- OpenAI API
- AI provider API key (OpenAI, Anthropic, Gemini, or other configured provider)
- Azure Document Intelligence
- Dropbox API
- Nextcloud instance
+3 -2
View File
@@ -26,7 +26,7 @@ Settings are organized into logical categories for easy navigation:
- **Core**: Database, Redis, working directory, external hostname, debug mode
- **Authentication**: Login settings, session secrets, OAuth configuration
- **AI Services**: OpenAI and Azure AI configuration
- **AI Services**: AI provider selection and credentials (OpenAI, Azure, Anthropic, Gemini, Ollama, OpenRouter, Portkey, LiteLLM)
- **Storage Providers**: Dropbox, Google Drive, OneDrive, S3, FTP, SFTP, WebDAV, Nextcloud, Paperless
- **Email**: SMTP configuration for sending emails
- **IMAP**: Email ingestion configuration (supports multiple accounts)
@@ -152,7 +152,8 @@ Removes a setting from the database (reverts to environment variable or default)
POST /api/settings/bulk-update
[
{"key": "debug", "value": "true"},
{"key": "openai_model", "value": "gpt-4"}
{"key": "ai_provider", "value": "anthropic"},
{"key": "openai_model", "value": "claude-3-5-sonnet-20241022"}
]
```
+1 -1
View File
@@ -150,7 +150,7 @@ View the complete processing history with a timeline showing:
- Timestamps for each operation
**Retry Processing**: If a file's processing has failed, you can use the "Retry Processing" button to reprocess the entire file. This is useful when:
- External API services (like OpenAI) had temporary issues
- External API services (like the configured AI provider) had temporary issues
- Network connectivity was lost during processing
- Configuration has been updated and you want to reprocess with new settings
+3 -2
View File
@@ -18,7 +18,8 @@
for everyone, whether you're a small startup or a large enterprise.
</p>
<p class="text-gray-600">
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.
</p>
@@ -33,7 +34,7 @@
<ul class="list-disc list-inside text-gray-600 ml-2">
<li>Simple and secure file uploads with drag & drop support</li>
<li>OCR powered by Azure Document Intelligence</li>
<li>Automated metadata extraction using OpenAI</li>
<li>Automated metadata extraction using a pluggable AI provider</li>
<li>PDF conversion for various file formats via Gotenberg</li>
<li>Intelligent document classification and date extraction</li>
</ul>
+14 -1
View File
@@ -139,8 +139,21 @@
Enable {{ setting.key.replace('_', ' ').title() }}
</label>
</div>
{% elif setting.metadata.options %}
<!-- Dropdown Select for fields with a fixed list of values -->
<select
id="{{ setting.key }}"
name="{{ setting.key }}"
x-model="formData['{{ setting.key }}']"
class="setting-input w-full px-3 py-2 border border-gray-300 rounded-md shadow-sm focus:outline-none focus:ring-blue-500 focus:border-blue-500"
>
<option value="">— select —</option>
{% for opt in setting.metadata.options %}
<option value="{{ opt }}">{{ opt }}</option>
{% endfor %}
</select>
{% else %}
<!-- Text Input -->
<!-- Text / Password Input -->
<div class="relative">
{% if setting.metadata.sensitive %}
<!-- Sensitive Field with Show/Hide Toggle -->
+15 -3
View File
@@ -92,7 +92,7 @@
<div class="border-b border-gray-200 pb-6 last:border-b-0">
<label for="{{ setting.key }}" class="block text-sm font-medium text-gray-900 mb-1">
{{ setting.label }}
{% if setting.default is none or setting.key in ['admin_password', 'openai_api_key', 'azure_ai_key', 'azure_endpoint'] %}
{% if setting.default is none or setting.key in ['admin_password'] %}
<span class="text-red-600">*</span>
{% endif %}
</label>
@@ -129,7 +129,18 @@
/>
</div>
{% else %}
<!-- Regular input -->
<!-- Regular input or select -->
{% if setting.options %}
<select
id="{{ setting.key }}"
name="{{ setting.key }}"
class="wizard-input w-full px-4 py-3 border border-gray-300 rounded-md shadow-sm focus:outline-none focus:ring-2 focus:ring-indigo-500 focus:border-transparent"
>
{% for opt in setting.options %}
<option value="{{ opt }}" {% if setting.current_value == opt %}selected{% elif not setting.current_value and opt == setting.default %}selected{% endif %}>{{ opt }}</option>
{% endfor %}
</select>
{% else %}
<input
type="{% if setting.sensitive %}password{% else %}text{% endif %}"
id="{{ setting.key }}"
@@ -137,9 +148,10 @@
value="{{ setting.current_value if setting.current_value else '' }}"
class="wizard-input w-full px-4 py-3 border border-gray-300 rounded-md shadow-sm focus:outline-none focus:ring-2 focus:ring-indigo-500 focus:border-transparent"
placeholder="{{ setting.description }}"
{% if setting.default is none or setting.key in ['admin_password', 'openai_api_key', 'azure_ai_key', 'azure_endpoint'] %}required{% endif %}
{% if setting.default is none or setting.key in ['admin_password'] %}required{% endif %}
/>
{% endif %}
{% endif %}
{% if setting.value_source == 'db' %}
<span class="inline-flex items-center px-2 py-0.5 rounded text-xs font-medium bg-green-100 text-green-800 mt-1">DB</span>
+4 -4
View File
@@ -174,11 +174,11 @@
Configure Now
</a>
{% endif %}
{% elif name == "OpenAI" %}
{% elif name == "AI Provider" %}
{% if provider.configured %}
<button
class="test-provider-btn inline-flex items-center px-2.5 py-1.5 border border-gray-300 text-xs font-medium rounded text-gray-700 bg-white hover:bg-gray-50 focus:outline-none focus:ring-2 focus:ring-offset-2 focus:ring-indigo-500"
data-provider="openai">
data-provider="ai_provider">
Test Connection
</button>
{% endif %}
@@ -521,8 +521,8 @@ document.addEventListener('DOMContentLoaded', function() {
endpoint = '/api/onedrive/test-token';
} else if (provider === 'google_drive') {
endpoint = '/api/google-drive/test-token';
} else if (provider === 'openai') {
endpoint = '/api/openai/test';
} else if (provider === 'ai_provider') {
endpoint = '/api/ai/test';
} else if (provider === 'azure') {
endpoint = '/api/azure/test';
}
+4 -1
View File
@@ -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
+775
View File
@@ -0,0 +1,775 @@
"""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,
PortkeyProvider,
_require_text_content,
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
def _make_none_content_response() -> MagicMock:
"""Build a mock response where message.content is None (e.g. tool call)."""
msg = SimpleNamespace(content=None)
choice = SimpleNamespace(message=msg)
resp = MagicMock()
resp.choices = [choice]
return resp
# ---------------------------------------------------------------------------
# _require_text_content helper
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestRequireTextContent:
"""Tests for the _require_text_content guard helper."""
def test_returns_non_none_string_unchanged(self):
"""Returns the string as-is when content is not None."""
assert _require_text_content("hello") == "hello"
def test_returns_empty_string_unchanged(self):
"""Returns an empty string as-is (empty ≠ None)."""
assert _require_text_content("") == ""
def test_raises_value_error_when_none(self):
"""Raises ValueError with descriptive message when content is None."""
with pytest.raises(ValueError, match="content=None"):
_require_text_content(None)
# ---------------------------------------------------------------------------
# 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
@patch("openai.OpenAI")
def test_chat_completion_raises_on_none_content(self, mock_openai_cls):
"""chat_completion raises ValueError when the response content is None."""
mock_client = MagicMock()
mock_client.chat.completions.create.return_value = _make_none_content_response()
mock_openai_cls.return_value = mock_client
provider = OpenAIProvider(api_key="sk-test")
with pytest.raises(ValueError, match="content=None"):
provider.chat_completion(
messages=[{"role": "user", "content": "test"}],
model="gpt-4o-mini",
)
# ---------------------------------------------------------------------------
# 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"
# ---------------------------------------------------------------------------
# PortkeyProvider
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestPortkeyProvider:
"""Tests for PortkeyProvider."""
@patch("openai.OpenAI")
def test_initialises_with_required_api_key(self, mock_openai_cls):
"""Provider passes x-portkey-api-key header and uses default gateway URL."""
PortkeyProvider(api_key="pk-key")
call_kwargs = mock_openai_cls.call_args[1]
assert call_kwargs["base_url"] == "https://api.portkey.ai/v1"
assert call_kwargs["api_key"] == "pk-key"
assert call_kwargs["default_headers"]["x-portkey-api-key"] == "pk-key"
@patch("openai.OpenAI")
def test_sets_virtual_key_header_when_provided(self, mock_openai_cls):
"""Provider includes x-portkey-virtual-key header when virtual_key is set."""
PortkeyProvider(api_key="pk-key", virtual_key="vk-abc123")
call_kwargs = mock_openai_cls.call_args[1]
assert call_kwargs["default_headers"]["x-portkey-virtual-key"] == "vk-abc123"
@patch("openai.OpenAI")
def test_omits_virtual_key_header_when_not_provided(self, mock_openai_cls):
"""Provider does not include x-portkey-virtual-key header when virtual_key is None."""
PortkeyProvider(api_key="pk-key")
call_kwargs = mock_openai_cls.call_args[1]
assert "x-portkey-virtual-key" not in call_kwargs["default_headers"]
@patch("openai.OpenAI")
def test_sets_config_header_when_provided(self, mock_openai_cls):
"""Provider includes x-portkey-config header when config is set."""
PortkeyProvider(api_key="pk-key", config="pc-my-config-xyz")
call_kwargs = mock_openai_cls.call_args[1]
assert call_kwargs["default_headers"]["x-portkey-config"] == "pc-my-config-xyz"
@patch("openai.OpenAI")
def test_omits_config_header_when_not_provided(self, mock_openai_cls):
"""Provider does not include x-portkey-config header when config is None."""
PortkeyProvider(api_key="pk-key")
call_kwargs = mock_openai_cls.call_args[1]
assert "x-portkey-config" not in call_kwargs["default_headers"]
@patch("openai.OpenAI")
def test_uses_custom_base_url(self, mock_openai_cls):
"""Provider forwards a custom gateway base URL."""
PortkeyProvider(api_key="pk-key", base_url="https://my-portkey-instance.example.com/v1")
call_kwargs = mock_openai_cls.call_args[1]
assert call_kwargs["base_url"] == "https://my-portkey-instance.example.com/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("Portkey response")
mock_openai_cls.return_value = mock_client
provider = PortkeyProvider(api_key="pk-key")
result = provider.chat_completion(
messages=[{"role": "user", "content": "hello"}],
model="gpt-4o",
)
assert result == "Portkey response"
@patch("openai.OpenAI")
def test_chat_completion_with_all_options(self, mock_openai_cls):
"""Provider works correctly with virtual_key, config, and custom model."""
mock_client = MagicMock()
mock_client.chat.completions.create.return_value = _make_openai_response("ok")
mock_openai_cls.return_value = mock_client
provider = PortkeyProvider(
api_key="pk-key",
virtual_key="vk-anthropic",
config="pc-fallback-config",
)
result = provider.chat_completion(
messages=[{"role": "user", "content": "test"}],
model="claude-3-5-sonnet-20241022",
temperature=0.5,
)
assert result == "ok"
call_kwargs = mock_client.chat.completions.create.call_args[1]
assert call_kwargs["model"] == "claude-3-5-sonnet-20241022"
assert call_kwargs["temperature"] == 0.5
# ---------------------------------------------------------------------------
# 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",
"portkey_api_key": None,
"portkey_virtual_key": None,
"portkey_config": None,
"portkey_base_url": "https://api.portkey.ai/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()
@patch("openai.OpenAI")
def test_returns_portkey_provider(self, mock_openai_cls):
"""get_ai_provider returns PortkeyProvider for ai_provider='portkey'."""
with patch(
"app.utils.ai_provider.settings",
self._mock_settings(
ai_provider="portkey",
portkey_api_key="pk-test",
portkey_virtual_key=None,
portkey_config=None,
portkey_base_url="https://api.portkey.ai/v1",
),
):
provider = get_ai_provider()
assert isinstance(provider, PortkeyProvider)
def test_portkey_raises_without_api_key(self):
"""get_ai_provider raises ValueError when portkey_api_key is missing."""
with patch(
"app.utils.ai_provider.settings",
self._mock_settings(
ai_provider="portkey",
portkey_api_key=None,
portkey_virtual_key=None,
portkey_config=None,
portkey_base_url="https://api.portkey.ai/v1",
),
):
with pytest.raises(ValueError, match="PORTKEY_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
+29 -29
View File
@@ -205,7 +205,7 @@ class TestCheckCredentialsTask:
@patch("app.tasks.check_credentials.get_failure_state")
@patch("app.tasks.check_credentials.get_provider_status")
@patch("app.tasks.check_credentials.validate_storage_configs")
@patch("app.tasks.check_credentials.sync_test_openai_connection")
@patch("app.tasks.check_credentials.sync_test_ai_provider_connection")
@patch("app.tasks.check_credentials.sync_test_azure_connection")
@patch("app.tasks.check_credentials.sync_test_dropbox_token")
@patch("app.tasks.check_credentials.sync_test_google_drive_token")
@@ -216,7 +216,7 @@ class TestCheckCredentialsTask:
mock_gdrive,
mock_dropbox,
mock_azure,
mock_openai,
mock_ai_provider,
mock_storage_configs,
mock_provider_status,
mock_get_state,
@@ -225,7 +225,7 @@ class TestCheckCredentialsTask:
"""Test checks all configured services."""
mock_get_state.return_value = {}
mock_provider_status.return_value = {
"OpenAI": {"configured": True},
"AI Provider": {"configured": True},
"Azure AI": {"configured": True},
"Dropbox": {"configured": True},
"Google Drive": {"configured": True},
@@ -234,7 +234,7 @@ class TestCheckCredentialsTask:
mock_storage_configs.return_value = {"dropbox": [], "google_drive": [], "onedrive": []}
# All tests succeed
mock_openai.return_value = {"status": "success"}
mock_ai_provider.return_value = {"status": "success"}
mock_azure.return_value = {"status": "success"}
mock_dropbox.return_value = {"status": "success"}
mock_gdrive.return_value = {"status": "success"}
@@ -249,14 +249,14 @@ class TestCheckCredentialsTask:
@patch("app.tasks.check_credentials.get_failure_state")
@patch("app.tasks.check_credentials.get_provider_status")
@patch("app.tasks.check_credentials.validate_storage_configs")
@patch("app.tasks.check_credentials.sync_test_openai_connection")
@patch("app.tasks.check_credentials.sync_test_ai_provider_connection")
def test_tracks_failures(
self, mock_openai, mock_storage_configs, mock_provider_status, mock_get_state, mock_save_state
self, mock_ai_provider, mock_storage_configs, mock_provider_status, mock_get_state, mock_save_state
):
"""Test tracks credential failures."""
mock_get_state.return_value = {}
mock_provider_status.return_value = {
"OpenAI": {"configured": True},
"AI Provider": {"configured": True},
"Azure AI": {"configured": False},
"Dropbox": {"configured": False},
"Google Drive": {"configured": False},
@@ -264,7 +264,7 @@ class TestCheckCredentialsTask:
}
mock_storage_configs.return_value = {}
mock_openai.return_value = {"status": "error", "message": "Invalid API key"}
mock_ai_provider.return_value = {"status": "error", "message": "Invalid API key"}
result = check_credentials()
@@ -281,7 +281,7 @@ class TestCheckCredentialsTask:
"""Test skips unconfigured services."""
mock_get_state.return_value = {}
mock_provider_status.return_value = {
"OpenAI": {"configured": False},
"AI Provider": {"configured": False},
"Azure AI": {"configured": False},
"Dropbox": {"configured": False},
"Google Drive": {"configured": False},
@@ -298,15 +298,15 @@ class TestCheckCredentialsTask:
@patch("app.tasks.check_credentials.get_failure_state")
@patch("app.tasks.check_credentials.get_provider_status")
@patch("app.tasks.check_credentials.validate_storage_configs")
@patch("app.tasks.check_credentials.sync_test_openai_connection")
@patch("app.tasks.check_credentials.sync_test_ai_provider_connection")
@patch("app.tasks.check_credentials.notify_credential_failure")
def test_sends_notifications_on_failure(
self, mock_notify, mock_openai, mock_storage_configs, mock_provider_status, mock_get_state, mock_save_state
self, mock_notify, mock_ai_provider, mock_storage_configs, mock_provider_status, mock_get_state, mock_save_state
):
"""Test sends notifications on credential failure."""
mock_get_state.return_value = {}
mock_provider_status.return_value = {
"OpenAI": {"configured": True},
"AI Provider": {"configured": True},
"Azure AI": {"configured": False},
"Dropbox": {"configured": False},
"Google Drive": {"configured": False},
@@ -314,7 +314,7 @@ class TestCheckCredentialsTask:
}
mock_storage_configs.return_value = {}
mock_openai.return_value = {"status": "error", "message": "Invalid API key"}
mock_ai_provider.return_value = {"status": "error", "message": "Invalid API key"}
check_credentials()
@@ -324,16 +324,16 @@ class TestCheckCredentialsTask:
@patch("app.tasks.check_credentials.get_failure_state")
@patch("app.tasks.check_credentials.get_provider_status")
@patch("app.tasks.check_credentials.validate_storage_configs")
@patch("app.tasks.check_credentials.sync_test_openai_connection")
@patch("app.tasks.check_credentials.sync_test_ai_provider_connection")
@patch("app.tasks.check_credentials.notify_credential_failure")
def test_suppresses_notifications_after_threshold(
self, mock_notify, mock_openai, mock_storage_configs, mock_provider_status, mock_get_state, mock_save_state
self, mock_notify, mock_ai_provider, mock_storage_configs, mock_provider_status, mock_get_state, mock_save_state
):
"""Test suppresses notifications after failure threshold."""
# Existing state with 4 failures
mock_get_state.return_value = {"OpenAI": {"count": 4, "last_notified": 12345}}
mock_get_state.return_value = {"AI Provider": {"count": 4, "last_notified": 12345}}
mock_provider_status.return_value = {
"OpenAI": {"configured": True},
"AI Provider": {"configured": True},
"Azure AI": {"configured": False},
"Dropbox": {"configured": False},
"Google Drive": {"configured": False},
@@ -341,7 +341,7 @@ class TestCheckCredentialsTask:
}
mock_storage_configs.return_value = {}
mock_openai.return_value = {"status": "error", "message": "Invalid API key"}
mock_ai_provider.return_value = {"status": "error", "message": "Invalid API key"}
check_credentials()
@@ -352,15 +352,15 @@ class TestCheckCredentialsTask:
@patch("app.tasks.check_credentials.get_failure_state")
@patch("app.tasks.check_credentials.get_provider_status")
@patch("app.tasks.check_credentials.validate_storage_configs")
@patch("app.tasks.check_credentials.sync_test_openai_connection")
@patch("app.tasks.check_credentials.sync_test_ai_provider_connection")
def test_tracks_recovery(
self, mock_openai, mock_storage_configs, mock_provider_status, mock_get_state, mock_save_state
self, mock_ai_provider, mock_storage_configs, mock_provider_status, mock_get_state, mock_save_state
):
"""Test tracks service recovery."""
# Existing state with failures
mock_get_state.return_value = {"OpenAI": {"count": 2, "last_notified": 12345}}
mock_get_state.return_value = {"AI Provider": {"count": 2, "last_notified": 12345}}
mock_provider_status.return_value = {
"OpenAI": {"configured": True},
"AI Provider": {"configured": True},
"Azure AI": {"configured": False},
"Dropbox": {"configured": False},
"Google Drive": {"configured": False},
@@ -369,7 +369,7 @@ class TestCheckCredentialsTask:
mock_storage_configs.return_value = {}
# Service is now valid
mock_openai.return_value = {"status": "success"}
mock_ai_provider.return_value = {"status": "success"}
result = check_credentials()
@@ -379,14 +379,14 @@ class TestCheckCredentialsTask:
@patch("app.tasks.check_credentials.get_failure_state")
@patch("app.tasks.check_credentials.get_provider_status")
@patch("app.tasks.check_credentials.validate_storage_configs")
@patch("app.tasks.check_credentials.sync_test_openai_connection")
@patch("app.tasks.check_credentials.sync_test_ai_provider_connection")
def test_handles_exception_during_check(
self, mock_openai, mock_storage_configs, mock_provider_status, mock_get_state, mock_save_state
self, mock_ai_provider, mock_storage_configs, mock_provider_status, mock_get_state, mock_save_state
):
"""Test handles exception during credential check."""
mock_get_state.return_value = {}
mock_provider_status.return_value = {
"OpenAI": {"configured": True},
"AI Provider": {"configured": True},
"Azure AI": {"configured": False},
"Dropbox": {"configured": False},
"Google Drive": {"configured": False},
@@ -394,11 +394,11 @@ class TestCheckCredentialsTask:
}
mock_storage_configs.return_value = {}
mock_openai.side_effect = Exception("Network error")
mock_ai_provider.side_effect = Exception("Network error")
result = check_credentials()
# Should still complete and record the error
assert result["failures"] == 1
assert "OpenAI" in result["results"]
assert result["results"]["OpenAI"]["status"] == "error"
assert "AI Provider" in result["results"]
assert result["results"]["AI Provider"]["status"] == "error"
+76 -81
View File
@@ -67,12 +67,11 @@ class TestExtractMetadataWithGpt:
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
@patch("app.tasks.extract_metadata_with_gpt.client")
def test_successful_metadata_extraction(self, mock_client, mock_log_progress, mock_embed_task):
"""Test successful metadata extraction with valid GPT response."""
# Mock the OpenAI client response
mock_completion = MagicMock()
mock_completion.choices[0].message.content = json.dumps(
@patch("app.tasks.extract_metadata_with_gpt.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.tasks.extract_metadata_with_gpt.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.tasks.extract_metadata_with_gpt.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.tasks.extract_metadata_with_gpt.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.tasks.extract_metadata_with_gpt.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.tasks.extract_metadata_with_gpt.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.tasks.extract_metadata_with_gpt.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.tasks.extract_metadata_with_gpt.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.tasks.extract_metadata_with_gpt.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.tasks.extract_metadata_with_gpt.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.tasks.extract_metadata_with_gpt.get_ai_provider")
@patch("app.tasks.extract_metadata_with_gpt.SessionLocal")
def test_database_lookup_with_existing_file(
self, mock_session_local, mock_client, mock_log_progress, mock_embed_task
self, mock_session_local, mock_get_provider, mock_log_progress, mock_embed_task
):
"""Test file_id retrieval when file exists on disk and in database (branches 76->82, 79->82)."""
mock_completion = MagicMock()
mock_completion.choices[0].message.content = '{"filename": "test.pdf", "document_type": "Unknown"}'
mock_client.chat.completions.create.return_value = mock_completion
"""Test file_id retrieval when file exists on disk and in database."""
mock_provider = MagicMock()
mock_provider.chat_completion.return_value = '{"filename": "test.pdf", "document_type": "Unknown"}'
mock_get_provider.return_value = mock_provider
# Mock database session
mock_db = MagicMock()
@@ -333,20 +334,14 @@ class TestExtractMetadataWithGpt:
@pytest.mark.unit
class TestClientInitialization:
"""Tests for OpenAI client initialization error handling."""
class TestModuleImports:
"""Tests for module import behaviour."""
def test_client_initialization_imports_successfully(self):
"""Test that module imports successfully even if client initialization fails (lines 25-27).
def test_module_imports_successfully(self):
"""Test that the module imports successfully."""
import app.tasks.extract_metadata_with_gpt as mod
The module has a try/except block for client initialization that sets client to None
on failure. This test verifies the module can be imported without crashing,
regardless of whether the client initializes successfully or not.
"""
# Import should succeed regardless of client initialization success
from app.tasks.extract_metadata_with_gpt import client
# Client will be either an OpenAI client instance or None
# Both are valid states - the important thing is the import doesn't crash
# We verify the client variable exists and has a defined type
assert hasattr(client, "__class__") or client is None
assert mod is not None
assert hasattr(mod, "extract_metadata_with_gpt")
assert hasattr(mod, "extract_json_from_text")
assert hasattr(mod, "get_ai_provider")
+18 -23
View File
@@ -407,65 +407,60 @@ startxref
@pytest.mark.unit
class TestRefineTextWithGPT:
"""Tests for OpenAI text refinement task."""
"""Tests for AI provider text refinement task."""
@patch("app.tasks.refine_text_with_gpt.log_task_progress")
def test_successful_text_refinement(self, mock_log):
"""Test successful text refinement with OpenAI."""
"""Test successful text refinement with AI provider."""
raw_text = "This is s0me text with OCR err0rs"
filename = "test.pdf"
# Mock OpenAI response
mock_choice = Mock()
mock_choice.message.content = "This is some text with OCR errors"
mock_response = Mock()
mock_response.choices = [mock_choice]
mock_client = Mock()
mock_client.chat.completions.create.return_value = mock_response
cleaned = "This is some text with OCR errors"
# Import the module to patch the correct function
from app.tasks import extract_metadata_with_gpt as metadata_module
mock_provider = MagicMock()
mock_provider.chat_completion.return_value = cleaned
with (
patch("app.tasks.refine_text_with_gpt.client", mock_client),
patch("app.tasks.refine_text_with_gpt.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.tasks.refine_text_with_gpt.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:
+1 -1
View File
@@ -35,7 +35,7 @@ class TestStatusDashboard:
mock_exists.return_value = False
mock_providers.return_value = {
"OpenAI": {"configured": True, "status": "success"},
"AI Provider": {"configured": True, "status": "success"},
"Azure AI": {"configured": False},
}
mock_settings.version = "1.0.0"