diff --git a/UPGRADE.md b/UPGRADE.md index 96d140f1a..810076cd1 100644 --- a/UPGRADE.md +++ b/UPGRADE.md @@ -33,6 +33,9 @@ Other changes: - The deprecated endpoint `/api/v1.0/documents//descendants` is removed. The search endpoint should be used instead. - Upgrade docspec dependency to version >= 3.0.0 The docspec service has changed since version 3.0.0, we ware now compatible with this version and not with version 2.x.x anymore +- It is now possible to use the Mistral SDK instead of the OpenAI for the AI features. If your provider is compatible with the mistral API, we encourage you to use it. +- `AI_API_KEY` settings is renamed in `OPENAI_SDK_API_KEY` and is only used to congiure the OpenAi sdk +- `AI_BASE_URL` settings is renamed in `OPENAI_SDK_BASE_URL` and is only used to congiure the OpenAi sdk ## [4.6.0] - 2026-02-27 diff --git a/src/backend/core/api/serializers.py b/src/backend/core/api/serializers.py index cef5adbca..5cfce9e95 100644 --- a/src/backend/core/api/serializers.py +++ b/src/backend/core/api/serializers.py @@ -19,7 +19,7 @@ from rest_framework import serializers from core import choices, enums, models, validators from core.services import mime_types -from core.services.ai_services import AI_ACTIONS +from core.services.ai_services.legacy import AI_ACTIONS from core.services.converter_services import ( ConversionError, Converter, diff --git a/src/backend/core/api/viewsets.py b/src/backend/core/api/viewsets.py index 8298f211f..f39c930e5 100644 --- a/src/backend/core/api/viewsets.py +++ b/src/backend/core/api/viewsets.py @@ -49,7 +49,8 @@ from treebeard.exceptions import InvalidMoveToDescendant from core import authentication, choices, enums, models from core.api.filters import remove_accents from core.services import mime_types -from core.services.ai_services import AIService +from core.services.ai_services.blocknote import AIService +from core.services.ai_services.legacy import get_legacy_ai_service from core.services.collaboration_services import CollaborationService from core.services.converter_services import ( ConversionError, @@ -2133,13 +2134,16 @@ class DocumentViewSet( # Check permissions first self.get_object() + if not settings.AI_FEATURE_ENABLED or not settings.AI_FEATURE_LEGACY_ENABLED: + raise ValidationError("AI feature is not enabled.") + serializer = serializers.AITransformSerializer(data=request.data) serializer.is_valid(raise_exception=True) text = serializer.validated_data["text"] action = serializer.validated_data["action"] - response = AIService().transform(text, action) + response = get_legacy_ai_service().transform(text, action) return drf.response.Response(response, status=drf.status.HTTP_200_OK) @@ -2161,13 +2165,16 @@ class DocumentViewSet( # Check permissions first self.get_object() + if not settings.AI_FEATURE_ENABLED or not settings.AI_FEATURE_LEGACY_ENABLED: + raise ValidationError("AI feature is not enabled.") + serializer = self.get_serializer(data=request.data) serializer.is_valid(raise_exception=True) text = serializer.validated_data["text"] language = serializer.validated_data["language"] - response = AIService().translate(text, language) + response = get_legacy_ai_service().translate(text, language) return drf.response.Response(response, status=drf.status.HTTP_200_OK) diff --git a/src/backend/core/services/ai_services.py b/src/backend/core/services/ai_services/blocknote.py similarity index 77% rename from src/backend/core/services/ai_services.py rename to src/backend/core/services/ai_services/blocknote.py index 0f6a59c68..d2b270ac9 100644 --- a/src/backend/core/services/ai_services.py +++ b/src/backend/core/services/ai_services/blocknote.py @@ -14,7 +14,6 @@ from django.conf import settings from django.core.exceptions import ImproperlyConfigured from langfuse import get_client -from langfuse.openai import OpenAI as OpenAI_Langfuse from pydantic_ai import Agent, DeferredToolRequests from pydantic_ai.models.mistral import MistralModel from pydantic_ai.models.openai import OpenAIChatModel @@ -27,13 +26,6 @@ from pydantic_ai.ui.vercel_ai import VercelAIAdapter from pydantic_ai.ui.vercel_ai.request_types import RequestData, TextUIPart, UIMessage from rest_framework.request import Request -from core import enums - -if settings.LANGFUSE_PUBLIC_KEY: - OpenAI = OpenAI_Langfuse -else: - from openai import OpenAI - log = logging.getLogger(__name__) BLOCKNOTE_TOOL_STRICT_PROMPT = """ @@ -67,50 +59,6 @@ IDs ALWAYS end with "$". Use ids EXACTLY as provided. Return ONLY the JSON tool input. No prose, no markdown. """ -AI_ACTIONS = { - "prompt": ( - "Answer the prompt using markdown formatting for structure and emphasis. " - "Return the content directly without wrapping it in code blocks or markdown delimiters. " - "Preserve the language and markdown formatting. " - "Do not provide any other information. " - "Preserve the language." - ), - "correct": ( - "Correct grammar and spelling of the markdown text, " - "preserving language and markdown formatting. " - "Do not provide any other information. " - "Preserve the language." - ), - "rephrase": ( - "Rephrase the given markdown text, " - "preserving language and markdown formatting. " - "Do not provide any other information. " - "Preserve the language." - ), - "summarize": ( - "Summarize the markdown text, preserving language and markdown formatting. " - "Do not provide any other information. " - "Preserve the language." - ), - "beautify": ( - "Add formatting to the text to make it more readable. " - "Do not provide any other information. " - "Preserve the language." - ), - "emojify": ( - "Add emojis to the important parts of the text. " - "Do not provide any other information. " - "Preserve the language." - ), -} - -AI_TRANSLATE = ( - "Keep the same html structure and formatting. " - "Translate the content in the html to the specified language {language:s}. " - "Check the translation for accuracy and make any necessary corrections. " - "Do not provide any other information." -) - def convert_async_generator_to_sync(async_gen: AsyncIterator[str]) -> Iterator[str]: """Convert an async generator to a sync generator.""" @@ -178,52 +126,9 @@ def configure_pydantic_model_provider() -> OpenAIChatModel | MistralModel: raise ImproperlyConfigured("AI configuration not set") -@cache -def configure_legacy_openai_client(): - """Configure the open ai sdk client for the legacy AI feature.""" - if ( - settings.OPENAI_SDK_BASE_URL is None - or settings.OPENAI_SDK_API_KEY is None - or settings.AI_MODEL is None - ): - raise ImproperlyConfigured("AI configuration not set") - return OpenAI( - base_url=settings.OPENAI_SDK_BASE_URL, api_key=settings.OPENAI_SDK_API_KEY - ) - - class AIService: """Service class for AI-related operations.""" - def call_ai_api(self, system_content, text): - """Helper method to call the OpenAI API and process the response.""" - client = configure_legacy_openai_client() - response = client.chat.completions.create( - model=settings.AI_MODEL, - messages=[ - {"role": "system", "content": system_content}, - {"role": "user", "content": text}, - ], - ) - - content = response.choices[0].message.content - - if not content: - raise RuntimeError("AI response does not contain an answer") - - return {"answer": content} - - def transform(self, text, action): - """Transform text based on specified action.""" - system_content = AI_ACTIONS[action] - return self.call_ai_api(system_content, text) - - def translate(self, text, language): - """Translate text to a specified language.""" - language_display = enums.ALL_LANGUAGES.get(language, language) - system_content = AI_TRANSLATE.format(language=language_display) - return self.call_ai_api(system_content, text) - @staticmethod def inject_document_state_messages( messages: list[UIMessage], diff --git a/src/backend/core/services/ai_services/legacy.py b/src/backend/core/services/ai_services/legacy.py new file mode 100644 index 000000000..19ee4193d --- /dev/null +++ b/src/backend/core/services/ai_services/legacy.py @@ -0,0 +1,178 @@ +"""Module dedicated to the legacy ai services.""" + +import logging +from abc import ABC, abstractmethod +from functools import cache + +from django.conf import settings +from django.core.exceptions import ImproperlyConfigured + +from langfuse.openai import OpenAI as OpenAI_Langfuse +from mistralai import Mistral + +from core import enums + +if settings.LANGFUSE_PUBLIC_KEY: + OpenAI = OpenAI_Langfuse +else: + from openai import OpenAI + +log = logging.getLogger(__name__) + +AI_ACTIONS = { + "prompt": ( + "Answer the prompt using markdown formatting for structure and emphasis. " + "Return the content directly without wrapping it in code blocks or markdown delimiters. " + "Preserve the language and markdown formatting. " + "Do not provide any other information. " + "Preserve the language." + ), + "correct": ( + "Correct grammar and spelling of the markdown text, " + "preserving language and markdown formatting. " + "Do not provide any other information. " + "Preserve the language." + ), + "rephrase": ( + "Rephrase the given markdown text, " + "preserving language and markdown formatting. " + "Do not provide any other information. " + "Preserve the language." + ), + "summarize": ( + "Summarize the markdown text, preserving language and markdown formatting. " + "Do not provide any other information. " + "Preserve the language." + ), + "beautify": ( + "Add formatting to the text to make it more readable. " + "Do not provide any other information. " + "Preserve the language." + ), + "emojify": ( + "Add emojis to the important parts of the text. " + "Do not provide any other information. " + "Preserve the language." + ), +} + +AI_TRANSLATE = ( + "Keep the same html structure and formatting. " + "Translate the content in the html to the specified language {language:s}. " + "Check the translation for accuracy and make any necessary corrections. " + "Do not provide any other information." +) + + +class LegacyAiClient(ABC): + """abstract class for legacy client.""" + + @abstractmethod + def call_ai_api(self, system_content, text) -> str: + """Abstract method call_ai_api.""" + + +class LegacyAiServiceMistralClient(LegacyAiClient): + """ai_service using mistral sdk for the legacy ai feature.""" + + def __init__(self): + """Configure mistral sdk""" + if ( + not settings.MISTRAL_SDK_API_KEY + or not settings.MISTRAL_SDK_BASE_URL + or not settings.AI_MODEL + ): + raise ImproperlyConfigured("Mistral sdk configuration not set") + + self.client = Mistral( + api_key=settings.MISTRAL_SDK_API_KEY, + server_url=settings.MISTRAL_SDK_BASE_URL, + ) + + def call_ai_api(self, system_content, text) -> str: + response = self.client.chat.complete( + model=settings.AI_MODEL, + messages=[ + {"role": "system", "content": system_content}, + {"role": "user", "content": text}, + ], + stream=False, + ) + + return response.choices[0].message.content + + +class LegacyAiServiceOpenAiClient(LegacyAiClient): + """ai_service using OpenAI client for the legacy ai feature.""" + + def __init__(self): + """configure OpenAI client.""" + if ( + not settings.OPENAI_SDK_BASE_URL + or not settings.OPENAI_SDK_API_KEY + or not settings.AI_MODEL + ): + raise ImproperlyConfigured("OpenAI configuration not set") + self.client = OpenAI( + base_url=settings.OPENAI_SDK_BASE_URL, api_key=settings.OPENAI_SDK_API_KEY + ) + + def call_ai_api(self, system_content, text) -> str: + response = self.client.chat.completions.create( + model=settings.AI_MODEL, + messages=[ + {"role": "system", "content": system_content}, + {"role": "user", "content": text}, + ], + ) + + return response.choices[0].message.content + + +class LegacyAIService: + """Legacy ai service used by transform and translate actions.""" + + def __init__(self, ai_client: LegacyAiClient): + """Assign client to the service.""" + self.ai_client = ai_client + + def call_ai_api(self, system_content, text): + """Helper method to call the OpenAI API and process the response.""" + + content = self.ai_client.call_ai_api(system_content, text) + + if not content: + raise RuntimeError("AI response does not contain an answer") + + return {"answer": content} + + def transform(self, text, action): + """Transform text based on specified action.""" + system_content = AI_ACTIONS[action] + return self.call_ai_api(system_content, text) + + def translate(self, text, language): + """Translate text to a specified language.""" + language_display = enums.ALL_LANGUAGES.get(language, language) + system_content = AI_TRANSLATE.format(language=language_display) + return self.call_ai_api(system_content, text) + + +@cache +def get_legacy_ai_service() -> LegacyAIService: + """Helper responsible to correctly instantiate and configure legacy ai service.""" + + ai_client = None + + if settings.MISTRAL_SDK_API_KEY: + ai_client = LegacyAiServiceMistralClient() + + if settings.OPENAI_SDK_API_KEY: + ai_client = LegacyAiServiceOpenAiClient() + + if not ai_client: + raise ImproperlyConfigured( + "trying to configure legacy ai_service but missing client configuration." + ) + + return LegacyAIService(ai_client) diff --git a/src/backend/core/tests/documents/test_api_documents_ai_proxy.py b/src/backend/core/tests/documents/test_api_documents_ai_proxy.py index ffda9d9ec..e99270f76 100644 --- a/src/backend/core/tests/documents/test_api_documents_ai_proxy.py +++ b/src/backend/core/tests/documents/test_api_documents_ai_proxy.py @@ -11,7 +11,7 @@ import pytest from rest_framework.test import APIClient from core import factories -from core.services.ai_services import configure_pydantic_model_provider +from core.services.ai_services.blocknote import configure_pydantic_model_provider from core.tests.conftest import TEAM, USER, VIA pytestmark = pytest.mark.django_db @@ -28,7 +28,6 @@ def ai_settings(settings): settings.AI_FEATURE_LEGACY_ENABLED = True settings.LANGFUSE_PUBLIC_KEY = None settings.AI_VERCEL_SDK_VERSION = 6 - yield configure_pydantic_model_provider.cache_clear() @@ -68,7 +67,7 @@ def test_api_documents_ai_proxy_anonymous_forbidden(reach, role): @override_settings(AI_ALLOW_REACH_FROM="public") -@patch("core.services.ai_services.AIService.stream") +@patch("core.services.ai_services.blocknote.AIService.stream") def test_api_documents_ai_proxy_anonymous_success(mock_stream): """ Anonymous users should be able to request AI proxy to a document @@ -152,7 +151,7 @@ def test_api_documents_ai_proxy_authenticated_forbidden(reach, role): ("public", "editor"), ], ) -@patch("core.services.ai_services.AIService.stream") +@patch("core.services.ai_services.blocknote.AIService.stream") def test_api_documents_ai_proxy_authenticated_success(mock_stream, reach, role): """ Authenticated users should be able to request AI proxy to a document @@ -208,7 +207,7 @@ def test_api_documents_ai_proxy_reader(via, mock_user_teams): @pytest.mark.parametrize("role", ["editor", "administrator", "owner"]) @pytest.mark.parametrize("via", VIA) -@patch("core.services.ai_services.AIService.stream") +@patch("core.services.ai_services.blocknote.AIService.stream") def test_api_documents_ai_proxy_success(mock_stream, via, role, mock_user_teams): """Users with sufficient permissions should be able to request AI proxy.""" user = factories.UserFactory() @@ -269,7 +268,7 @@ def test_api_documents_ai_proxy_ai_feature_disabled(settings, setting_to_disable @override_settings(AI_DOCUMENT_RATE_THROTTLE_RATES={"minute": 3, "hour": 6, "day": 10}) -@patch("core.services.ai_services.AIService.stream") +@patch("core.services.ai_services.blocknote.AIService.stream") def test_api_documents_ai_proxy_throttling_document(mock_stream): """ Throttling per document should be triggered on the AI proxy endpoint. @@ -307,7 +306,7 @@ def test_api_documents_ai_proxy_throttling_document(mock_stream): @override_settings(AI_USER_RATE_THROTTLE_RATES={"minute": 3, "hour": 6, "day": 10}) -@patch("core.services.ai_services.AIService.stream") +@patch("core.services.ai_services.blocknote.AIService.stream") def test_api_documents_ai_proxy_throttling_user(mock_stream): """ Throttling per user should be triggered on the AI proxy endpoint. @@ -342,7 +341,7 @@ def test_api_documents_ai_proxy_throttling_user(mock_stream): } -@patch("core.services.ai_services.AIService.stream") +@patch("core.services.ai_services.blocknote.AIService.stream") def test_api_documents_ai_proxy_returns_streaming_response(mock_stream): """AI proxy should return a StreamingHttpResponse with correct headers.""" user = factories.UserFactory() diff --git a/src/backend/core/tests/documents/test_api_documents_ai_transform.py b/src/backend/core/tests/documents/test_api_documents_ai_transform.py index 5ada04c3c..f42847c8a 100644 --- a/src/backend/core/tests/documents/test_api_documents_ai_transform.py +++ b/src/backend/core/tests/documents/test_api_documents_ai_transform.py @@ -8,7 +8,7 @@ import pytest from rest_framework.test import APIClient from core import factories -from core.services.ai_services import configure_legacy_openai_client +from core.services.ai_services.legacy import get_legacy_ai_service from core.tests.conftest import TEAM, USER, VIA pytestmark = pytest.mark.django_db @@ -17,6 +17,8 @@ pytestmark = pytest.mark.django_db @pytest.fixture def ai_settings(settings): """Fixture to set AI settings.""" + settings.AI_FEATURE_ENABLED = True + settings.AI_FEATURE_LEGACY_ENABLED = True settings.OPENAI_SDK_BASE_URL = "http://example.com" settings.OPENAI_SDK_API_KEY = "test-key" settings.AI_MODEL = "llama" @@ -25,8 +27,7 @@ def ai_settings(settings): @pytest.fixture(autouse=True) def clear_openai_client_config(): """Clear the _configure_legacy_openai_client cache""" - yield - configure_legacy_openai_client.cache_clear() + get_legacy_ai_service.cache_clear() @pytest.mark.parametrize( @@ -37,7 +38,7 @@ def clear_openai_client_config(): ("restricted", "reader", "restricted"), ("restricted", "editor", "public"), ("restricted", "editor", "authenticated"), - ("restricted", "editor", "restrictied"), + ("restricted", "editor", "restricted"), ("authenticated", "reader", "public"), ("authenticated", "reader", "authenticated"), ("authenticated", "reader", "restricted"), @@ -281,6 +282,7 @@ def test_api_documents_ai_transform_success(mock_create, via, role, mock_user_te ) +@pytest.mark.usefixtures("ai_settings") def test_api_documents_ai_transform_empty_text(): """The text should not be empty when requesting AI transform.""" user = factories.UserFactory() @@ -297,6 +299,7 @@ def test_api_documents_ai_transform_empty_text(): assert response.json() == {"text": ["This field may not be blank."]} +@pytest.mark.usefixtures("ai_settings") def test_api_documents_ai_transform_invalid_action(): """The action should valid when requesting AI transform.""" user = factories.UserFactory() diff --git a/src/backend/core/tests/documents/test_api_documents_ai_translate.py b/src/backend/core/tests/documents/test_api_documents_ai_translate.py index c800a588e..1954cbc47 100644 --- a/src/backend/core/tests/documents/test_api_documents_ai_translate.py +++ b/src/backend/core/tests/documents/test_api_documents_ai_translate.py @@ -8,7 +8,7 @@ import pytest from rest_framework.test import APIClient from core import factories -from core.services.ai_services import configure_legacy_openai_client +from core.services.ai_services.legacy import get_legacy_ai_service from core.tests.conftest import TEAM, USER, VIA pytestmark = pytest.mark.django_db @@ -17,6 +17,8 @@ pytestmark = pytest.mark.django_db @pytest.fixture def ai_settings(settings): """Fixture to set AI settings.""" + settings.AI_FEATURE_ENABLED = True + settings.AI_FEATURE_LEGACY_ENABLED = True settings.OPENAI_SDK_BASE_URL = "http://example.com" settings.OPENAI_SDK_API_KEY = "test-key" settings.AI_MODEL = "llama" @@ -25,8 +27,7 @@ def ai_settings(settings): @pytest.fixture(autouse=True) def clear_openai_client_config(): "clear the configure_legacy_openai_client cache" - yield - configure_legacy_openai_client.cache_clear() + get_legacy_ai_service.cache_clear() def test_api_documents_ai_translate_viewset_options_metadata(): @@ -57,7 +58,7 @@ def test_api_documents_ai_translate_viewset_options_metadata(): ("restricted", "reader", "restricted"), ("restricted", "editor", "public"), ("restricted", "editor", "authenticated"), - ("restricted", "editor", "restrictied"), + ("restricted", "editor", "restricted"), ("authenticated", "reader", "public"), ("authenticated", "reader", "authenticated"), ("authenticated", "reader", "restricted"), @@ -303,6 +304,7 @@ def test_api_documents_ai_translate_success(mock_create, via, role, mock_user_te ) +@pytest.mark.usefixtures("ai_settings") def test_api_documents_ai_translate_empty_text(): """The text should not be empty when requesting AI translate.""" user = factories.UserFactory() @@ -319,6 +321,7 @@ def test_api_documents_ai_translate_empty_text(): assert response.json() == {"text": ["This field may not be blank."]} +@pytest.mark.usefixtures("ai_settings") def test_api_documents_ai_translate_invalid_action(): """The action should valid when requesting AI translate.""" user = factories.UserFactory() diff --git a/src/backend/core/tests/external_api/test_external_api_documents_ai.py b/src/backend/core/tests/external_api/test_external_api_documents_ai.py index 92480fe81..ee695958b 100644 --- a/src/backend/core/tests/external_api/test_external_api_documents_ai.py +++ b/src/backend/core/tests/external_api/test_external_api_documents_ai.py @@ -14,7 +14,7 @@ import pytest from rest_framework.test import APIClient from core import factories, models -from core.services.ai_services import configure_legacy_openai_client +from core.services.ai_services.legacy import get_legacy_ai_service from core.tests.documents.test_api_documents_ai_proxy import ( # pylint: disable=unused-import ai_settings, ) @@ -27,8 +27,7 @@ pytestmark = pytest.mark.django_db @pytest.fixture(autouse=True) def clear_openai_client_config(): """Clear the configure_legacy_openai_client cache.""" - yield - configure_legacy_openai_client.cache_clear() + get_legacy_ai_service.cache_clear() def test_external_api_documents_ai_transform_not_allowed( @@ -249,7 +248,7 @@ def test_external_api_documents_ai_translate_can_be_allowed( } ) @pytest.mark.usefixtures("ai_settings") -@patch("core.services.ai_services.AIService.stream") +@patch("core.services.ai_services.blocknote.AIService.stream") def test_external_api_documents_ai_proxy_can_be_allowed( mock_stream, user_token, resource_server_backend, user_specific_sub ): diff --git a/src/backend/core/tests/test_services_ai_services.py b/src/backend/core/tests/test_services_ai_services.py index bcdba86c8..3d61f6f71 100644 --- a/src/backend/core/tests/test_services_ai_services.py +++ b/src/backend/core/tests/test_services_ai_services.py @@ -10,18 +10,23 @@ from django.core.exceptions import ImproperlyConfigured from django.test.utils import override_settings import pytest +from mistralai import Mistral from openai import OpenAI, OpenAIError from pydantic_ai.models.mistral import MistralModel from pydantic_ai.models.openai import OpenAIChatModel from pydantic_ai.ui.vercel_ai.request_types import TextUIPart, UIMessage -from core.services.ai_services import ( +from core.services.ai_services.blocknote import ( BLOCKNOTE_TOOL_STRICT_PROMPT, AIService, - configure_legacy_openai_client, configure_pydantic_model_provider, convert_async_generator_to_sync, ) +from core.services.ai_services.legacy import ( + LegacyAiServiceMistralClient, + LegacyAiServiceOpenAiClient, + get_legacy_ai_service, +) pytestmark = pytest.mark.django_db @@ -39,7 +44,7 @@ def ai_settings(settings): settings.AI_VERCEL_SDK_VERSION = 6 yield configure_pydantic_model_provider.cache_clear() - configure_legacy_openai_client.cache_clear() + get_legacy_ai_service.cache_clear() # -- AIService configure sdk-- @@ -53,7 +58,7 @@ def ai_settings(settings): ("AI_MODEL", None), ], ) -def test_ai_services_configure_legacy_openai_sdk_missing( +def test_ai_services_configure_open_ai_leagcy_client_missing_settings( setting_name, setting_value, settings ): """ @@ -65,18 +70,57 @@ def test_ai_services_configure_legacy_openai_sdk_missing( ImproperlyConfigured, match="AI configuration not set", ): - configure_legacy_openai_client() + LegacyAiServiceOpenAiClient() -def test_ai_services_configure_legacy_openai_sdk(settings): - """With all required settings an open ai sdk instance should be configured.""" +def test_ai_services_configure_open_ai_leagcy_client(settings): + """With all required settings the OpenAi legacy client should be configured.""" settings.AI_MODEL = "llama" settings.OPENAI_SDK_BASE_URL = "http://example.com" settings.OPENAI_SDK_API_KEY = "test-key" - openai_sdk = configure_legacy_openai_client() + legacy_openai_client = LegacyAiServiceOpenAiClient() - assert isinstance(openai_sdk, OpenAI) + assert isinstance(legacy_openai_client.client, OpenAI) + + +@pytest.mark.parametrize( + "setting_name, setting_value", + [ + ("MISTRAL_SDK_BASE_URL", None), + ("MISTRAL_SDK_API_KEY", None), + ("AI_MODEL", None), + ], +) +def test_ai_services_configure_mistral_sdk_leagcy_client_missing_settings( + setting_name, setting_value, settings +): + """ + An exception must be raised if an expected settings is missing to configure the openai sdk. + """ + settings.OPENAI_SDK_BASE_URL = None + settings.OPENAI_SDK_API_KEY = None + setattr(settings, setting_name, setting_value) + + with pytest.raises( + ImproperlyConfigured, + match="Mistral sdk configuration not set", + ): + LegacyAiServiceMistralClient() + + +def test_ai_services_configure_mistral_sdk_legacy_client(settings): + """With all required settings the Mistral sdk legacy client should be configured.""" + + settings.AI_MODEL = "llama" + settings.OPENAI_SDK_BASE_URL = None + settings.OPENAI_SDK_API_KEY = None + settings.MISTRAL_SDK_API_KEY = "mistreal-sdk-key" + settings.MISTRAL_SDK_BASE_URL = "https://mistral.base-url.com" + + legacy_mistral_client = LegacyAiServiceMistralClient() + + assert isinstance(legacy_mistral_client.client, Mistral) def test_ai_services_configure_pydantic_ai_model_openai(settings): @@ -132,7 +176,7 @@ def test_services_ai_client_error(mock_create): OpenAIError, match="Mocked client error", ): - AIService().transform("hello", "prompt") + get_legacy_ai_service().transform("hello", "prompt") @override_settings( @@ -150,7 +194,7 @@ def test_services_ai_client_invalid_response(mock_create): RuntimeError, match="AI response does not contain an answer", ): - AIService().transform("hello", "prompt") + get_legacy_ai_service().transform("hello", "prompt") @override_settings( @@ -164,7 +208,7 @@ def test_services_ai_success(mock_create): choices=[MagicMock(message=MagicMock(content="Salut"))] ) - response = AIService().transform("hello", "prompt") + response = get_legacy_ai_service().transform("hello", "prompt") assert response == {"answer": "Salut"} @@ -180,7 +224,7 @@ def test_services_ai_translate_success(mock_create): choices=[MagicMock(message=MagicMock(content="Bonjour"))] ) - response = AIService().translate("

Hello

", "fr") + response = get_legacy_ai_service().translate("

Hello

", "fr") assert response == {"answer": "Bonjour"} call_args = mock_create.call_args @@ -196,7 +240,7 @@ def test_services_ai_translate_unknown_language(mock_create): choices=[MagicMock(message=MagicMock(content="Translated"))] ) - response = AIService().translate("

Hello

", "xx-unknown") + response = get_legacy_ai_service().translate("

Hello

", "xx-unknown") assert response == {"answer": "Translated"} call_args = mock_create.call_args @@ -507,7 +551,7 @@ def test_services_ai_stream_defaults_to_sync(mock_build, monkeypatch): # -- AIService._build_async_stream -- -@patch("core.services.ai_services.VercelAIAdapter") +@patch("core.services.ai_services.blocknote.VercelAIAdapter") def test_services_ai_build_async_stream(mock_adapter_cls): """_build_async_stream should build the pydantic-ai streaming pipeline.""" @@ -536,7 +580,7 @@ def test_services_ai_build_async_stream(mock_adapter_cls): mock_adapter_instance.encode_stream.assert_called_once() -@patch("core.services.ai_services.VercelAIAdapter") +@patch("core.services.ai_services.blocknote.VercelAIAdapter") def test_services_ai_build_async_stream_with_tool_definitions(mock_adapter_cls): """_build_async_stream should build an ExternalToolset when toolDefinitions are present in the request.""" @@ -573,7 +617,7 @@ def test_services_ai_build_async_stream_with_tool_definitions(mock_adapter_cls): assert len(call_kwargs["toolsets"]) == 1 -@patch("core.services.ai_services.VercelAIAdapter") +@patch("core.services.ai_services.blocknote.VercelAIAdapter") def test_services_ai_build_async_stream_with_tool_definitions_required_system_prompt( mock_adapter_cls, ): @@ -616,8 +660,8 @@ def test_services_ai_build_async_stream_with_tool_definitions_required_system_pr assert mock_run_input.messages[0].parts[0].text == BLOCKNOTE_TOOL_STRICT_PROMPT -@patch("core.services.ai_services.Agent") -@patch("core.services.ai_services.VercelAIAdapter") +@patch("core.services.ai_services.blocknote.Agent") +@patch("core.services.ai_services.blocknote.VercelAIAdapter") def test_services_ai_build_async_stream_langfuse_enabled( mock_adapter_cls, mock_agent_cls, settings ):