From 6cfc8990b96498993dc2350affab7bd3247168ea Mon Sep 17 00:00:00 2001 From: Manuel Raynaud Date: Tue, 28 Apr 2026 10:42:39 +0200 Subject: [PATCH] =?UTF-8?q?=E2=99=BB=EF=B8=8F(backend)=20use=20mistral=20s?= =?UTF-8?q?dk=20with=20legacy=20ai=20feature?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit We also want to use the mistral sdk with the legacy AI feature when this one is configured with the settings. In order to separate bot feature, they now live in their own module. --- UPGRADE.md | 3 + src/backend/core/api/serializers.py | 2 +- src/backend/core/api/viewsets.py | 13 +- .../blocknote.py} | 95 ---------- .../core/services/ai_services/legacy.py | 178 ++++++++++++++++++ .../documents/test_api_documents_ai_proxy.py | 15 +- .../test_api_documents_ai_transform.py | 11 +- .../test_api_documents_ai_translate.py | 11 +- .../test_external_api_documents_ai.py | 7 +- .../core/tests/test_services_ai_services.py | 82 ++++++-- 10 files changed, 279 insertions(+), 138 deletions(-) rename src/backend/core/services/{ai_services.py => ai_services/blocknote.py} (77%) create mode 100644 src/backend/core/services/ai_services/legacy.py 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 ):