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:
@@ -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
@@ -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
@@ -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/) |
|
||||||
|
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
Reference in New Issue
Block a user