feat: add Portkey provider support and null-content guard to AI abstraction layer

Co-authored-by: christianlouis <361235+christianlouis@users.noreply.github.com>
This commit is contained in:
copilot-swe-agent[bot]
2026-02-23 19:42:03 +00:00
parent d4c7fb26ac
commit dcdef44303
6 changed files with 457 additions and 27 deletions
+7
View File
@@ -35,6 +35,13 @@ class Settings(BaseSettings):
openrouter_api_key: Optional[str] = None openrouter_api_key: Optional[str] = None
openrouter_base_url: str = "https://openrouter.ai/api/v1" 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 (used when ai_provider="azure")
azure_openai_api_version: str = "2024-02-01" azure_openai_api_version: str = "2024-02-01"
workdir: str workdir: str
+119 -12
View File
@@ -3,8 +3,8 @@
This module provides a pluggable abstraction for various AI model providers, This module provides a pluggable abstraction for various AI model providers,
allowing the platform to work with OpenAI, Azure OpenAI, Anthropic Claude, allowing the platform to work with OpenAI, Azure OpenAI, Anthropic Claude,
Google Gemini, Ollama (local LLMs), OpenRouter, and any LiteLLM-compatible Google Gemini, Ollama (local LLMs), OpenRouter, Portkey, and any
provider without being locked to a single vendor. LiteLLM-compatible provider without being locked to a single vendor.
Provider selection is controlled by the ``AI_PROVIDER`` environment variable. Provider selection is controlled by the ``AI_PROVIDER`` environment variable.
See the Configuration Guide for full details on each provider's settings. See the Configuration Guide for full details on each provider's settings.
@@ -14,9 +14,36 @@ import logging
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import Any, Dict, List, Optional from typing import Any, Dict, List, Optional
from app.config import settings
logger = logging.getLogger(__name__) 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): class AIProvider(ABC):
"""Abstract base class for AI chat completion providers. """Abstract base class for AI chat completion providers.
@@ -80,7 +107,8 @@ class OpenAIProvider(AIProvider):
temperature=temperature, temperature=temperature,
**kwargs, **kwargs,
) )
return completion.choices[0].message.content _content = completion.choices[0].message.content
return _require_text_content(_content)
class AzureOpenAIProvider(AIProvider): class AzureOpenAIProvider(AIProvider):
@@ -108,7 +136,8 @@ class AzureOpenAIProvider(AIProvider):
temperature=temperature, temperature=temperature,
**kwargs, **kwargs,
) )
return completion.choices[0].message.content _content = completion.choices[0].message.content
return _require_text_content(_content)
class AnthropicProvider(AIProvider): class AnthropicProvider(AIProvider):
@@ -139,7 +168,8 @@ class AnthropicProvider(AIProvider):
api_key=self._api_key, api_key=self._api_key,
**kwargs, **kwargs,
) )
return response.choices[0].message.content _content = response.choices[0].message.content
return _require_text_content(_content)
class GeminiProvider(AIProvider): class GeminiProvider(AIProvider):
@@ -170,7 +200,8 @@ class GeminiProvider(AIProvider):
api_key=self._api_key, api_key=self._api_key,
**kwargs, **kwargs,
) )
return response.choices[0].message.content _content = response.choices[0].message.content
return _require_text_content(_content)
class OllamaProvider(AIProvider): class OllamaProvider(AIProvider):
@@ -210,7 +241,8 @@ class OllamaProvider(AIProvider):
temperature=temperature, temperature=temperature,
**kwargs, **kwargs,
) )
return completion.choices[0].message.content _content = completion.choices[0].message.content
return _require_text_content(_content)
class OpenRouterProvider(AIProvider): class OpenRouterProvider(AIProvider):
@@ -243,7 +275,74 @@ class OpenRouterProvider(AIProvider):
temperature=temperature, temperature=temperature,
**kwargs, **kwargs,
) )
return completion.choices[0].message.content _content = completion.choices[0].message.content
return _require_text_content(_content)
class PortkeyProvider(AIProvider):
"""Portkey AI gateway (https://portkey.ai).
Portkey is an AI gateway that provides observability, caching, automatic
retries, fallbacks, and load balancing across 200+ LLMs via a single
OpenAI-compatible endpoint.
Required settings:
``PORTKEY_API_KEY`` your Portkey account API key.
Optional settings:
``PORTKEY_VIRTUAL_KEY`` a Portkey *Virtual Key* that maps to the
credentials of a specific provider stored in your Portkey vault.
When set, you do not need to expose the underlying provider's API key
in your environment.
``PORTKEY_CONFIG`` a saved Portkey *Config* ID (e.g.
``pc-my-config-abc123``) that applies advanced routing rules such as
fallbacks and load balancing.
``PORTKEY_BASE_URL`` override the gateway endpoint.
Default: ``https://api.portkey.ai/v1``.
The model name should match what the underlying provider expects (e.g.
``gpt-4o`` for OpenAI, ``claude-3-5-sonnet-20241022`` for Anthropic via a
virtual key).
"""
def __init__(
self,
api_key: str,
virtual_key: Optional[str] = None,
config: Optional[str] = None,
base_url: str = "https://api.portkey.ai/v1",
) -> None:
import openai
portkey_headers: Dict[str, str] = {"x-portkey-api-key": api_key}
if virtual_key:
portkey_headers["x-portkey-virtual-key"] = virtual_key
if config:
portkey_headers["x-portkey-config"] = config
self._client = openai.OpenAI(
api_key=api_key,
base_url=base_url,
default_headers=portkey_headers,
)
def chat_completion(
self,
messages: List[Dict[str, str]],
model: str,
temperature: float = 0,
**kwargs: Any,
) -> str:
completion = self._client.chat.completions.create(
model=model,
messages=messages,
temperature=temperature,
**kwargs,
)
_content = completion.choices[0].message.content
return _require_text_content(_content)
class LiteLLMProvider(AIProvider): class LiteLLMProvider(AIProvider):
@@ -283,7 +382,8 @@ class LiteLLMProvider(AIProvider):
completion_kwargs["api_base"] = self._api_base completion_kwargs["api_base"] = self._api_base
completion_kwargs.update(kwargs) completion_kwargs.update(kwargs)
response = litellm.completion(**completion_kwargs) response = litellm.completion(**completion_kwargs)
return response.choices[0].message.content _content = response.choices[0].message.content
return _require_text_content(_content)
def get_ai_provider() -> AIProvider: def get_ai_provider() -> AIProvider:
@@ -300,8 +400,6 @@ def get_ai_provider() -> AIProvider:
ValueError: If the configured provider name is not recognised. ValueError: If the configured provider name is not recognised.
ValueError: If required credentials for the selected provider are absent. ValueError: If required credentials for the selected provider are absent.
""" """
from app.config import settings
provider = settings.ai_provider.lower() provider = settings.ai_provider.lower()
logger.debug(f"Creating AI provider: {provider}") logger.debug(f"Creating AI provider: {provider}")
@@ -333,6 +431,15 @@ def get_ai_provider() -> AIProvider:
api_key=settings.openrouter_api_key, api_key=settings.openrouter_api_key,
base_url=settings.openrouter_base_url, 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": elif provider == "litellm":
return LiteLLMProvider( return LiteLLMProvider(
api_key=settings.openai_api_key or None, api_key=settings.openai_api_key or None,
@@ -341,5 +448,5 @@ def get_ai_provider() -> AIProvider:
else: else:
raise ValueError( raise ValueError(
f"Unknown AI provider: '{provider}'. " f"Unknown AI provider: '{provider}'. "
"Supported providers: openai, azure, anthropic, gemini, ollama, openrouter, litellm" "Supported providers: openai, azure, anthropic, gemini, ollama, openrouter, portkey, litellm"
) )
+145 -2
View File
@@ -302,11 +302,154 @@ SECURITY_HEADER_CSP_VALUE="default-src 'self'; script-src 'self' https://trusted
- [Deployment Guide - Security Headers](DeploymentGuide.md#security-headers) for Traefik/Nginx examples - [Deployment Guide - Security Headers](DeploymentGuide.md#security-headers) for Traefik/Nginx examples
- [SECURITY_AUDIT.md](../SECURITY_AUDIT.md#infrastructure-security) for security rationale - [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** | | **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_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/) | | `AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT` | Endpoint URL for Azure Doc Intelligence API. | [Azure Portal](https://portal.azure.com/) |
+173
View File
@@ -14,6 +14,8 @@ from app.utils.ai_provider import (
OllamaProvider, OllamaProvider,
OpenAIProvider, OpenAIProvider,
OpenRouterProvider, OpenRouterProvider,
PortkeyProvider,
_require_text_content,
get_ai_provider, get_ai_provider,
) )
@@ -40,6 +42,37 @@ def _make_litellm_response(content: str) -> MagicMock:
return resp 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 # Abstract base class
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -129,6 +162,20 @@ class TestOpenAIProvider:
assert call_kwargs["max_tokens"] == 100 assert call_kwargs["max_tokens"] == 100
assert call_kwargs["temperature"] == 0.7 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 # AzureOpenAIProvider
@@ -343,6 +390,97 @@ class TestOpenRouterProvider:
assert result == "Router response" 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 # LiteLLMProvider
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -444,6 +582,10 @@ class TestGetAIProvider:
"ollama_base_url": "http://localhost:11434", "ollama_base_url": "http://localhost:11434",
"openrouter_api_key": None, "openrouter_api_key": None,
"openrouter_base_url": "https://openrouter.ai/api/v1", "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) defaults.update(kwargs)
mock_settings = MagicMock() mock_settings = MagicMock()
@@ -536,6 +678,37 @@ class TestGetAIProvider:
with pytest.raises(ValueError, match="OPENROUTER_API_KEY"): with pytest.raises(ValueError, match="OPENROUTER_API_KEY"):
get_ai_provider() 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): def test_returns_litellm_provider(self):
"""get_ai_provider returns LiteLLMProvider for ai_provider='litellm'.""" """get_ai_provider returns LiteLLMProvider for ai_provider='litellm'."""
with patch( with patch(
+11 -11
View File
@@ -67,7 +67,7 @@ class TestExtractMetadataWithGpt:
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf") @patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
@patch("app.utils.ai_provider.get_ai_provider") @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): def test_successful_metadata_extraction(self, mock_get_provider, mock_log_progress, mock_embed_task):
"""Test successful metadata extraction with valid AI provider response.""" """Test successful metadata extraction with valid AI provider response."""
mock_provider = MagicMock() mock_provider = MagicMock()
@@ -116,7 +116,7 @@ class TestExtractMetadataWithGpt:
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf") @patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
@patch("app.utils.ai_provider.get_ai_provider") @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): def test_handles_json_in_backticks(self, mock_get_provider, mock_log_progress, mock_embed_task):
"""Test extraction handles JSON wrapped in markdown code blocks.""" """Test extraction handles JSON wrapped in markdown code blocks."""
mock_provider = MagicMock() mock_provider = MagicMock()
@@ -135,7 +135,7 @@ class TestExtractMetadataWithGpt:
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf") @patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
@patch("app.utils.ai_provider.get_ai_provider") @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): 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.""" """Test handling of invalid JSON in AI provider response."""
mock_provider = MagicMock() mock_provider = MagicMock()
@@ -154,7 +154,7 @@ class TestExtractMetadataWithGpt:
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf") @patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
@patch("app.utils.ai_provider.get_ai_provider") @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): def test_handles_api_exception(self, mock_get_provider, mock_log_progress, mock_embed_task):
"""Test handling of AI provider API exceptions.""" """Test handling of AI provider API exceptions."""
mock_provider = MagicMock() mock_provider = MagicMock()
@@ -173,7 +173,7 @@ class TestExtractMetadataWithGpt:
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf") @patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
@patch("app.utils.ai_provider.get_ai_provider") @patch("app.tasks.extract_metadata_with_gpt.get_ai_provider")
@patch("app.tasks.extract_metadata_with_gpt.SessionLocal") @patch("app.tasks.extract_metadata_with_gpt.SessionLocal")
def test_retrieves_file_id_from_database_when_not_provided( def test_retrieves_file_id_from_database_when_not_provided(
self, mock_session_local, mock_get_provider, mock_log_progress, mock_embed_task self, mock_session_local, mock_get_provider, mock_log_progress, mock_embed_task
@@ -207,7 +207,7 @@ class TestExtractMetadataWithGpt:
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf") @patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
@patch("app.utils.ai_provider.get_ai_provider") @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): def test_validates_filename_security(self, mock_get_provider, mock_log_progress, mock_embed_task):
"""Test filename validation to prevent path traversal.""" """Test filename validation to prevent path traversal."""
mock_provider = MagicMock() mock_provider = MagicMock()
@@ -227,7 +227,7 @@ class TestExtractMetadataWithGpt:
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf") @patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
@patch("app.utils.ai_provider.get_ai_provider") @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): def test_validates_filename_with_dots(self, mock_get_provider, mock_log_progress, mock_embed_task):
"""Test filename validation rejects '..' in filenames.""" """Test filename validation rejects '..' in filenames."""
mock_provider = MagicMock() mock_provider = MagicMock()
@@ -245,7 +245,7 @@ class TestExtractMetadataWithGpt:
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf") @patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
@patch("app.utils.ai_provider.get_ai_provider") @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): def test_accepts_valid_filename(self, mock_get_provider, mock_log_progress, mock_embed_task):
"""Test that valid filenames are accepted.""" """Test that valid filenames are accepted."""
mock_provider = MagicMock() mock_provider = MagicMock()
@@ -263,7 +263,7 @@ class TestExtractMetadataWithGpt:
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf") @patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
@patch("app.utils.ai_provider.get_ai_provider") @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): def test_handles_malformed_json_with_valid_structure(self, mock_get_provider, mock_log_progress, mock_embed_task):
"""Test handling of JSON that's parseable but missing expected fields.""" """Test handling of JSON that's parseable but missing expected fields."""
mock_provider = MagicMock() mock_provider = MagicMock()
@@ -280,7 +280,7 @@ class TestExtractMetadataWithGpt:
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf") @patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
@patch("app.utils.ai_provider.get_ai_provider") @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): 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.""" """Test handling when filename is provided as an absolute path."""
mock_provider = MagicMock() mock_provider = MagicMock()
@@ -299,7 +299,7 @@ class TestExtractMetadataWithGpt:
@patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf") @patch("app.tasks.extract_metadata_with_gpt.embed_metadata_into_pdf")
@patch("app.tasks.extract_metadata_with_gpt.log_task_progress") @patch("app.tasks.extract_metadata_with_gpt.log_task_progress")
@patch("app.utils.ai_provider.get_ai_provider") @patch("app.tasks.extract_metadata_with_gpt.get_ai_provider")
@patch("app.tasks.extract_metadata_with_gpt.SessionLocal") @patch("app.tasks.extract_metadata_with_gpt.SessionLocal")
def test_database_lookup_with_existing_file( def test_database_lookup_with_existing_file(
self, mock_session_local, mock_get_provider, mock_log_progress, mock_embed_task self, mock_session_local, mock_get_provider, mock_log_progress, mock_embed_task
+2 -2
View File
@@ -423,7 +423,7 @@ class TestRefineTextWithGPT:
mock_provider.chat_completion.return_value = cleaned mock_provider.chat_completion.return_value = cleaned
with ( with (
patch("app.utils.ai_provider.get_ai_provider", return_value=mock_provider), 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.object(metadata_module, "extract_metadata_with_gpt") as mock_extract,
patch("app.tasks.refine_text_with_gpt.settings") as mock_settings, patch("app.tasks.refine_text_with_gpt.settings") as mock_settings,
): ):
@@ -456,7 +456,7 @@ class TestRefineTextWithGPT:
mock_provider.chat_completion.side_effect = Exception("OpenAI API error") mock_provider.chat_completion.side_effect = Exception("OpenAI API error")
with ( with (
patch("app.utils.ai_provider.get_ai_provider", return_value=mock_provider), 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, patch("app.tasks.refine_text_with_gpt.settings") as mock_settings,
): ):
mock_settings.openai_model = "gpt-4" mock_settings.openai_model = "gpt-4"