diff --git a/app/config.py b/app/config.py index f9a18df5..702d6485 100644 --- a/app/config.py +++ b/app/config.py @@ -35,6 +35,13 @@ class Settings(BaseSettings): openrouter_api_key: Optional[str] = None openrouter_base_url: str = "https://openrouter.ai/api/v1" + # Portkey AI gateway settings (used when ai_provider="portkey") + # See https://portkey.ai for setup instructions + portkey_api_key: Optional[str] = None + portkey_virtual_key: Optional[str] = None # Routes to a specific provider via Portkey vault + portkey_config: Optional[str] = None # Portkey Config ID for advanced routing rules + portkey_base_url: str = "https://api.portkey.ai/v1" + # Azure OpenAI API version (used when ai_provider="azure") azure_openai_api_version: str = "2024-02-01" workdir: str diff --git a/app/utils/ai_provider.py b/app/utils/ai_provider.py index a98cae90..fa9353b0 100644 --- a/app/utils/ai_provider.py +++ b/app/utils/ai_provider.py @@ -3,8 +3,8 @@ This module provides a pluggable abstraction for various AI model providers, allowing the platform to work with OpenAI, Azure OpenAI, Anthropic Claude, -Google Gemini, Ollama (local LLMs), OpenRouter, and any LiteLLM-compatible -provider without being locked to a single vendor. +Google Gemini, Ollama (local LLMs), OpenRouter, Portkey, and any +LiteLLM-compatible provider without being locked to a single vendor. Provider selection is controlled by the ``AI_PROVIDER`` environment variable. See the Configuration Guide for full details on each provider's settings. @@ -14,9 +14,36 @@ import logging from abc import ABC, abstractmethod from typing import Any, Dict, List, Optional +from app.config import settings + logger = logging.getLogger(__name__) +def _require_text_content(content: Optional[str]) -> str: + """Raise a clear error if the AI response contains no text content. + + This can happen when the model returns a tool/function call instead of + a plain text message. All DocuElevate prompts expect a plain-text or + JSON response, so ``None`` content is always an unexpected condition. + + Args: + content: The ``message.content`` value from the completion response. + + Returns: + The original string, guaranteed non-None. + + Raises: + ValueError: If *content* is ``None``. + """ + if content is None: + raise ValueError( + "AI provider returned a response with no text content (content=None). " + "This may occur when the model generates a tool call instead of a plain text reply. " + "Ensure the model and prompt are configured for text/JSON output." + ) + return content + + class AIProvider(ABC): """Abstract base class for AI chat completion providers. @@ -80,7 +107,8 @@ class OpenAIProvider(AIProvider): temperature=temperature, **kwargs, ) - return completion.choices[0].message.content + _content = completion.choices[0].message.content + return _require_text_content(_content) class AzureOpenAIProvider(AIProvider): @@ -108,7 +136,8 @@ class AzureOpenAIProvider(AIProvider): temperature=temperature, **kwargs, ) - return completion.choices[0].message.content + _content = completion.choices[0].message.content + return _require_text_content(_content) class AnthropicProvider(AIProvider): @@ -139,7 +168,8 @@ class AnthropicProvider(AIProvider): api_key=self._api_key, **kwargs, ) - return response.choices[0].message.content + _content = response.choices[0].message.content + return _require_text_content(_content) class GeminiProvider(AIProvider): @@ -170,7 +200,8 @@ class GeminiProvider(AIProvider): api_key=self._api_key, **kwargs, ) - return response.choices[0].message.content + _content = response.choices[0].message.content + return _require_text_content(_content) class OllamaProvider(AIProvider): @@ -210,7 +241,8 @@ class OllamaProvider(AIProvider): temperature=temperature, **kwargs, ) - return completion.choices[0].message.content + _content = completion.choices[0].message.content + return _require_text_content(_content) class OpenRouterProvider(AIProvider): @@ -243,7 +275,74 @@ class OpenRouterProvider(AIProvider): temperature=temperature, **kwargs, ) - return completion.choices[0].message.content + _content = completion.choices[0].message.content + return _require_text_content(_content) + + +class PortkeyProvider(AIProvider): + """Portkey AI gateway (https://portkey.ai). + + Portkey is an AI gateway that provides observability, caching, automatic + retries, fallbacks, and load balancing across 200+ LLMs via a single + OpenAI-compatible endpoint. + + Required settings: + ``PORTKEY_API_KEY`` – your Portkey account API key. + + Optional settings: + ``PORTKEY_VIRTUAL_KEY`` – a Portkey *Virtual Key* that maps to the + credentials of a specific provider stored in your Portkey vault. + When set, you do not need to expose the underlying provider's API key + in your environment. + + ``PORTKEY_CONFIG`` – a saved Portkey *Config* ID (e.g. + ``pc-my-config-abc123``) that applies advanced routing rules such as + fallbacks and load balancing. + + ``PORTKEY_BASE_URL`` – override the gateway endpoint. + Default: ``https://api.portkey.ai/v1``. + + The model name should match what the underlying provider expects (e.g. + ``gpt-4o`` for OpenAI, ``claude-3-5-sonnet-20241022`` for Anthropic via a + virtual key). + """ + + def __init__( + self, + api_key: str, + virtual_key: Optional[str] = None, + config: Optional[str] = None, + base_url: str = "https://api.portkey.ai/v1", + ) -> None: + import openai + + portkey_headers: Dict[str, str] = {"x-portkey-api-key": api_key} + if virtual_key: + portkey_headers["x-portkey-virtual-key"] = virtual_key + if config: + portkey_headers["x-portkey-config"] = config + + self._client = openai.OpenAI( + api_key=api_key, + base_url=base_url, + default_headers=portkey_headers, + ) + + def chat_completion( + self, + messages: List[Dict[str, str]], + model: str, + temperature: float = 0, + **kwargs: Any, + ) -> str: + completion = self._client.chat.completions.create( + model=model, + messages=messages, + temperature=temperature, + **kwargs, + ) + _content = completion.choices[0].message.content + return _require_text_content(_content) class LiteLLMProvider(AIProvider): @@ -283,7 +382,8 @@ class LiteLLMProvider(AIProvider): completion_kwargs["api_base"] = self._api_base completion_kwargs.update(kwargs) response = litellm.completion(**completion_kwargs) - return response.choices[0].message.content + _content = response.choices[0].message.content + return _require_text_content(_content) def get_ai_provider() -> AIProvider: @@ -300,8 +400,6 @@ def get_ai_provider() -> AIProvider: ValueError: If the configured provider name is not recognised. ValueError: If required credentials for the selected provider are absent. """ - from app.config import settings - provider = settings.ai_provider.lower() logger.debug(f"Creating AI provider: {provider}") @@ -333,6 +431,15 @@ def get_ai_provider() -> AIProvider: api_key=settings.openrouter_api_key, base_url=settings.openrouter_base_url, ) + elif provider == "portkey": + if not settings.portkey_api_key: + raise ValueError("PORTKEY_API_KEY must be set when AI_PROVIDER='portkey'") + return PortkeyProvider( + api_key=settings.portkey_api_key, + virtual_key=settings.portkey_virtual_key, + config=settings.portkey_config, + base_url=settings.portkey_base_url, + ) elif provider == "litellm": return LiteLLMProvider( api_key=settings.openai_api_key or None, @@ -341,5 +448,5 @@ def get_ai_provider() -> AIProvider: else: raise ValueError( f"Unknown AI provider: '{provider}'. " - "Supported providers: openai, azure, anthropic, gemini, ollama, openrouter, litellm" + "Supported providers: openai, azure, anthropic, gemini, ollama, openrouter, portkey, litellm" ) diff --git a/docs/ConfigurationGuide.md b/docs/ConfigurationGuide.md index de57935b..b44cbbae 100644 --- a/docs/ConfigurationGuide.md +++ b/docs/ConfigurationGuide.md @@ -302,11 +302,154 @@ SECURITY_HEADER_CSP_VALUE="default-src 'self'; script-src 'self' https://trusted - [Deployment Guide - Security Headers](DeploymentGuide.md#security-headers) for Traefik/Nginx examples - [SECURITY_AUDIT.md](../SECURITY_AUDIT.md#infrastructure-security) for security rationale -### OpenAI & Azure Document Intelligence +### AI Provider & Model Selection + +DocuElevate supports multiple AI providers for metadata extraction and OCR text refinement. Select the provider via `AI_PROVIDER` and configure the matching credentials below. + +| **Variable** | **Description** | **Default** | +|-------------------|-----------------------------------------------------------------------|--------------------| +| `AI_PROVIDER` | Active AI provider. See supported values below. | `openai` | +| `AI_MODEL` | Model name for the selected provider. Falls back to `OPENAI_MODEL` when not set. | *(unset)* | +| `OPENAI_MODEL` | Default model name (used when `AI_MODEL` is not set). | `gpt-4o-mini` | + +**Supported `AI_PROVIDER` values**: `openai`, `azure`, `anthropic`, `gemini`, `ollama`, `openrouter`, `portkey`, `litellm` + +--- + +#### OpenAI (default) + +| **Variable** | **Description** | **Default** | +|-----------------------|--------------------------------------------------|----------------------------------| +| `OPENAI_API_KEY` | OpenAI API key. | *(required)* | +| `OPENAI_BASE_URL` | API base URL. Change for compatible proxies. | `https://api.openai.com/v1` | + +```bash +AI_PROVIDER=openai +OPENAI_API_KEY=sk-... +OPENAI_MODEL=gpt-4o-mini +``` + +#### Azure OpenAI + +| **Variable** | **Description** | **Default** | +|-------------------------------|----------------------------------------------|----------------| +| `OPENAI_API_KEY` | Azure OpenAI API key. | *(required)* | +| `OPENAI_BASE_URL` | Azure resource endpoint URL. | *(required)* | +| `AZURE_OPENAI_API_VERSION` | Azure OpenAI API version string. | `2024-02-01` | + +```bash +AI_PROVIDER=azure +OPENAI_API_KEY= +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/) | diff --git a/tests/test_ai_provider.py b/tests/test_ai_provider.py index e77caefa..d957c550 100644 --- a/tests/test_ai_provider.py +++ b/tests/test_ai_provider.py @@ -14,6 +14,8 @@ from app.utils.ai_provider import ( OllamaProvider, OpenAIProvider, OpenRouterProvider, + PortkeyProvider, + _require_text_content, get_ai_provider, ) @@ -40,6 +42,37 @@ def _make_litellm_response(content: str) -> MagicMock: 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 # --------------------------------------------------------------------------- @@ -129,6 +162,20 @@ class TestOpenAIProvider: 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 @@ -343,6 +390,97 @@ class TestOpenRouterProvider: 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 # --------------------------------------------------------------------------- @@ -444,6 +582,10 @@ class TestGetAIProvider: "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() @@ -536,6 +678,37 @@ class TestGetAIProvider: 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( diff --git a/tests/test_extract_metadata_gpt.py b/tests/test_extract_metadata_gpt.py index ff96bad8..0fc8f39a 100644 --- a/tests/test_extract_metadata_gpt.py +++ b/tests/test_extract_metadata_gpt.py @@ -67,7 +67,7 @@ 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.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): """Test successful metadata extraction with valid AI provider response.""" 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.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): """Test extraction handles JSON wrapped in markdown code blocks.""" 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.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): """Test handling of invalid JSON in AI provider response.""" 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.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): """Test handling of AI provider API exceptions.""" 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.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") def test_retrieves_file_id_from_database_when_not_provided( 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.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): """Test filename validation to prevent path traversal.""" 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.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): """Test filename validation rejects '..' in filenames.""" 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.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): """Test that valid filenames are accepted.""" 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.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): """Test handling of JSON that's parseable but missing expected fields.""" 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.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): """Test handling when filename is provided as an absolute path.""" 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.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") def test_database_lookup_with_existing_file( self, mock_session_local, mock_get_provider, mock_log_progress, mock_embed_task diff --git a/tests/test_ocr_processing.py b/tests/test_ocr_processing.py index c2c9159e..abf10289 100644 --- a/tests/test_ocr_processing.py +++ b/tests/test_ocr_processing.py @@ -423,7 +423,7 @@ class TestRefineTextWithGPT: mock_provider.chat_completion.return_value = cleaned 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("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") 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, ): mock_settings.openai_model = "gpt-4"