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:
@@ -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
|
||||
|
||||
@@ -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
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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 {}
|
||||
|
||||
@@ -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",
|
||||
)
|
||||
|
||||
|
||||
@@ -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 (0–1). Default: 0 (deterministic).
|
||||
**kwargs: Additional provider-specific arguments passed through.
|
||||
|
||||
Returns:
|
||||
The model's response as a plain string.
|
||||
|
||||
Raises:
|
||||
Exception: If the underlying API call fails.
|
||||
"""
|
||||
|
||||
|
||||
class OpenAIProvider(AIProvider):
|
||||
"""OpenAI provider using the ``openai`` Python SDK.
|
||||
|
||||
Also works as a drop-in for any OpenAI-compatible API endpoint, including
|
||||
LocalAI and LM Studio. Ollama and OpenRouter have dedicated providers with
|
||||
sensible defaults, but this provider works for them too when a custom
|
||||
``base_url`` is supplied.
|
||||
"""
|
||||
|
||||
def __init__(self, api_key: str, base_url: Optional[str] = None) -> None:
|
||||
import openai
|
||||
|
||||
self._client = openai.OpenAI(
|
||||
api_key=api_key,
|
||||
base_url=base_url or "https://api.openai.com/v1",
|
||||
)
|
||||
|
||||
def chat_completion(
|
||||
self,
|
||||
messages: List[Dict[str, str]],
|
||||
model: str,
|
||||
temperature: float = 0,
|
||||
**kwargs: Any,
|
||||
) -> str:
|
||||
completion = self._client.chat.completions.create(
|
||||
model=model,
|
||||
messages=messages,
|
||||
temperature=temperature,
|
||||
**kwargs,
|
||||
)
|
||||
_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"
|
||||
)
|
||||
@@ -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
@@ -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
@@ -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
@@ -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/) |
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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>
|
||||
|
||||
@@ -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 -->
|
||||
|
||||
@@ -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>
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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"
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user