@@ -11,4 +11,4 @@ class AuthenticationConfig(AppConfig): verbose_name = 'Пользователи' def ready(self): - pass + from .signals import invalidate_user_cache @@ -61,21 +61,6 @@ class Flux(SimpleService): def make(self, input_message: Message, save: bool = True) -> list[Message]: start_time = time.time() - version = 'flux-schnell' - user_prompt = self.translate_prompt(input_message.content) - callback_data = dict( - { - 'prompt': f'{user_prompt}\n{self.OPTIMIZATION_PROMPT}', - 'go_fast': False, - 'output_quality': 100, - **input_message.info, - } - ) - runner = replicate_run( - f'{self._CALLBACK_BASE}{version}', - callback_data, - ) - images = runner if isinstance(runner, list) else [runner] output_megapixels = 1 callback_data = { 'prompt': input_message.content, @@ -112,6 +112,10 @@ class NeuronModel(BaseModel, OrderedModel): def versions(self) -> QuerySet['ModelVersion']: return self.model_modelversions.all() + @property + def default_version(self) -> 'ModelVersion | None': + return self.versions.order_by('order').first() + @property def payment_rules(self) -> QuerySet['ModelPaymentRule']: return self.model_modelpaymentrules.all() @@ -2,13 +2,13 @@ import base64 from io import BytesIO import filetype -from django.db.models import Q from django.core.files.uploadedfile import InMemoryUploadedFile from django.utils.translation import gettext as _ from ninja.errors import HttpError from ml_model.models import NeuronModel from tools.chats.schemas import MessageInSchema +from tools.public_api.services.defaults import resolve_model_in_category from tools.public_api.services.openai_errors import OpenAIErrorService from tools.public_api.services.openai_stream import OpenAIStreamService @@ -73,14 +73,7 @@ def _parse_body(body: dict) -> MessageInSchema: def _resolve_model(model_ref: str) -> NeuronModel: try: - return ( - NeuronModel.objects.filter( - Q(model_modelversions__slug=model_ref) | Q(slug=model_ref), - category__slug='chat-bots', - ) - .distinct() - .get() - ) + return resolve_model_in_category(model_ref, 'chat-bots') except NeuronModel.DoesNotExist: raise HttpError(404, _('Model not found')) @@ -7,8 +7,8 @@ from ninja.errors import HttpError from django.utils.translation import gettext as _ from messages.models import Message -from ml_model.exceptions import FileNotProvided, InvalidParameterError -from ml_model.selectors.ml_models_selector import NeuronModelSelector +from ml_model.choices import ContentTypes +from ml_model.exceptions import FileNotProvided, InvalidParameterError, NeuronModelNotExist from ml_model.validators import ModelInputValidator from tools.chats.schemas import MessageInSchema @@ -17,6 +17,7 @@ from tools.chats.services.sse_store import PublicSSEStoreService from tools.chats.tasks import public_event_stream_task from tools.public_api.models import APIKey, APIStore +from tools.public_api.services.defaults import resolve_model, resolve_version router = Router(auth=None, tags=['public']) @@ -100,20 +101,25 @@ def public_stream_message(request, model_slug: str, body: MessageInSchema): api_store, created = APIStore.objects.get_or_create(user=user) - selector = NeuronModelSelector(user) - model = selector.get_model_by_slug(slug=model_slug) + try: + model = resolve_model(model_slug, user, ContentTypes.TEXT) + except NeuronModelNotExist as exc: + raise HttpError(404, str(exc)) if model.blocked: raise HttpError(403, _('Model is blocked by outdating or temporary block, please retry later')) if not model.streaming: raise HttpError(501, _('Stream not supported for this model')) data = body.model_dump(include={'content', 'file', 'info'}, exclude_unset=True) + info = data.get('info') or {} + resolve_version(model, info) + data['info'] = info try: ModelInputValidator( model, content=data.get('content'), file=data.get('file'), - info=data.get('info'), + info=info, ).validate() except (FileNotProvided, InvalidParameterError) as exc: raise HttpError(400, str(exc)) @@ -128,7 +134,7 @@ def public_stream_message(request, model_slug: str, body: MessageInSchema): start_user_balance=balance, message_uuid=str(input_message.pk), user_uuid=str(user.pk), - model_slug=model_slug, + model_slug=model.slug, api_key_uuid=str(api_key.pk), debit_api_key_limit=api_key.token_limit is not None, ) @@ -0,0 +1,64 @@ +from django.db.models import Q + +from authentication.models.user import CustomUserModel +from ml_model.exceptions import NeuronModelNotExist +from ml_model.models import NeuronModel +from ml_model.selectors.ml_models_selector import NeuronModelSelector + +DEFAULT_SLUG = 'default' + + +def resolve_model(slug: str, user: CustomUserModel, content_type: str) -> NeuronModel: + selector = NeuronModelSelector(user) + if slug != DEFAULT_SLUG: + return selector.get_model_by_slug(slug=slug) + model = ( + selector.get_models_by_output_content_type(output_content_type=content_type, serialize=False) + .filter( + model_settings__isnull=False, + model_settings__is_active=True, + private_models_hosts__isnull=True, + ) + .order_by('order') + .first() + ) + if model is None: + raise NeuronModelNotExist + return model + + +def resolve_model_in_category(model_ref: str, category_slug: str) -> NeuronModel: + if model_ref == DEFAULT_SLUG: + model = ( + NeuronModel.objects.filter( + category__slug=category_slug, + model_settings__isnull=False, + model_settings__is_active=True, + private_models_hosts__isnull=True, + ) + .order_by('order') + .first() + ) + if model is None: + raise NeuronModel.DoesNotExist + return model + return ( + NeuronModel.objects.filter( + Q(model_modelversions__slug=model_ref) | Q(slug=model_ref), + category__slug=category_slug, + ) + .distinct() + .get() + ) + + +def resolve_version(model: NeuronModel, info: dict) -> dict: + version = info.get('version') + if version not in (None, '', DEFAULT_SLUG): + return info + first = model.default_version + if first: + info['version'] = first.slug + else: + info.pop('version', None) + return info @@ -0,0 +1,63 @@ +from django.test import TestCase + +from authentication.models.user import CustomUserModel +from ml_model.exceptions import NeuronModelNotExist +from ml_model.models import ModelCategory, ModelSettings, ModelVersion, NeuronModel +from payments.models import PaymentPlan +from tools.public_api.services.defaults import resolve_model, resolve_model_in_category, resolve_version + + +class PublicDefaultResolveTest(TestCase): + @classmethod + def setUpTestData(cls) -> None: + PaymentPlan.objects.update_or_create(price=0, tokens_per_plan=10, defaults={}) + cls.user = CustomUserModel.objects.create_user(email='default-test@test.test', password='test') + cls.category = ModelCategory.objects.create(title='Chat-bots', slug='chat-bots') + cls.inactive = NeuronModel.objects.create( + title='Inactive', slug='inactive-bot', category=cls.category + ) + cls.model = NeuronModel.objects.create(title='Active', slug='active-bot', category=cls.category) + ModelSettings.objects.create(model=cls.inactive, is_active=False) + ModelSettings.objects.create(model=cls.model, is_active=True) + cls.first_version = ModelVersion.objects.create(model=cls.model, name='First', slug='first-ver') + cls.second_version = ModelVersion.objects.create(model=cls.model, name='Second', slug='second-ver') + cls.plain = NeuronModel.objects.create(title='Plain', slug='plain-bot', category=cls.category) + ModelSettings.objects.create(model=cls.plain, is_active=True) + + def test_default_version_is_first_by_order(self) -> None: + info = {'version': 'default'} + resolve_version(self.model, info) + self.assertEqual(info['version'], self.first_version.slug) + self.assertLess(self.first_version.order, self.second_version.order) + + def test_missing_version_uses_first(self) -> None: + info: dict = {} + resolve_version(self.model, info) + self.assertEqual(info['version'], self.first_version.slug) + + def test_explicit_version_is_kept(self) -> None: + info = {'version': self.second_version.slug} + resolve_version(self.model, info) + self.assertEqual(info['version'], self.second_version.slug) + + def test_no_versions_drops_default(self) -> None: + info = {'version': 'default', 'foo': 1} + resolve_version(self.plain, info) + self.assertNotIn('version', info) + self.assertEqual(info['foo'], 1) + + def test_default_model_is_first_active_in_category(self) -> None: + model = resolve_model('default', self.user, 'text') + self.assertEqual(model, self.model) + + def test_named_slug_is_unchanged(self) -> None: + model = resolve_model(self.plain.slug, self.user, 'text') + self.assertEqual(model, self.plain) + + def test_default_model_in_category(self) -> None: + model = resolve_model_in_category('default', 'chat-bots') + self.assertEqual(model, self.model) + + def test_unknown_default_category_raises(self) -> None: + with self.assertRaises(NeuronModelNotExist): + resolve_model('default', self.user, 'image') @@ -2,13 +2,13 @@ import httpx from uuid import uuid4 from django.http import HttpResponse -from django.db.models import Q from django.utils.translation import gettext_lazy as _ from rest_framework import status from rest_framework.response import Response from ml_model.models import NeuronModel from tools.public_api.permissions import HasElevenlabsAPIKey +from tools.public_api.services.defaults import resolve_model_in_category from tools.public_api.views import VoiceView from tools.public_api.views.voice import ( PublicVoiceDetailAPIView, @@ -133,14 +133,7 @@ class ElevenlabsVoiceCloneAPIView(BaseElevenlabsAPIView, VoiceView): ) ) try: - model_slug = ( - NeuronModel.objects.filter( - Q(model_modelversions__slug=str(model_ref)) | Q(slug=str(model_ref)), - category__slug='voice', - ) - .get() - .slug - ) + model_slug = resolve_model_in_category(str(model_ref), 'voice').slug except NeuronModel.DoesNotExist: return _elevenlabs_error_response( Response({'detail': _('Model not found')}, status=status.HTTP_400_BAD_REQUEST) @@ -6,7 +6,6 @@ from io import BytesIO import filetype import httpx from django.core.files.uploadedfile import InMemoryUploadedFile -from django.db.models import Q from django.http import HttpResponse from django.utils.translation import gettext_lazy as _ from rest_framework import status @@ -16,6 +15,7 @@ from rest_framework.response import Response from ml_model.models import NeuronModel from tools.public_api.selectors.api_key import APIKeySelector +from tools.public_api.services.defaults import resolve_model_in_category from tools.public_api.views.base import BaseGenerationView from tools.public_api.views.ml_service import VoiceView from tools.public_api.views.voice import ( @@ -98,14 +98,7 @@ class OpenAICompatibleAPIView(BaseGenerationView): data['info']['version'] = version try: - model_slug = ( - NeuronModel.objects.filter( - Q(model_modelversions__slug=data['info']['version']) | Q(slug=data['info']['version']), - category__slug='chat-bots', - ) - .get() - .slug - ) + model_slug = resolve_model_in_category(data['info']['version'], 'chat-bots').slug except NeuronModel.DoesNotExist: return _openai_error(_('Model not found')) @@ -235,14 +228,7 @@ class OpenAIAudioSpeechAPIView(VoiceView): code='missing_required_parameter', ) try: - model_slug = ( - NeuronModel.objects.filter( - Q(model_modelversions__slug=model_ref) | Q(slug=model_ref), - category__slug='voice', - ) - .get() - .slug - ) + model_slug = resolve_model_in_category(model_ref, 'voice').slug except NeuronModel.DoesNotExist: return _openai_error(_('Model not found'), param='model') voice_raw = request.data.get('voice') @@ -7,6 +7,7 @@ from rest_framework.status import ( HTTP_400_BAD_REQUEST, HTTP_402_PAYMENT_REQUIRED, HTTP_403_FORBIDDEN, + HTTP_404_NOT_FOUND, HTTP_500_INTERNAL_SERVER_ERROR, ) from rest_framework.views import APIView @@ -26,6 +27,7 @@ from ml_model.exceptions import ( InputImageSensitiveContentError, InvalidParameterError, ModelVersionNotAvailable, + NeuronModelNotExist, OutputSensitiveImageContentError, PaidPlanRequiredError, PromptLengthExceeded, @@ -40,6 +42,7 @@ from payments.exceptions.insufficient_balance import InsufficientBalance from tools.public_api.models import APIKey, APIStore from tools.public_api.permissions import HasAPIKey from tools.public_api.selectors.api_key import APIKeySelector +from tools.public_api.services.defaults import resolve_model, resolve_version logger = logging.getLogger(__name__) @@ -80,7 +83,10 @@ class BaseGenerationView(APIView): if key.token_limit is not None and key.token_limit < 1: return Response({'detail': _('Key limit exceeded')}, HTTP_403_FORBIDDEN) store, created = APIStore.objects.get_or_create(user=user) - model: NeuronModel = NeuronModelSelector(store.user).get_model_by_slug(slug=model_slug) + try: + model: NeuronModel = resolve_model(model_slug, store.user, self.output_content_type) + except NeuronModelNotExist as exc: + return Response({'detail': str(exc)}, status=HTTP_404_NOT_FOUND) if model.blocked: return Response( {'detail': _('Model is blocked by outdating or temporary block, please retry later')}, @@ -89,7 +95,8 @@ class BaseGenerationView(APIView): serializer = MessageSerializer(data=request.data) serializer.is_valid(raise_exception=True) service = model.service - info = serializer.validated_data.pop('info', {}) + info = serializer.validated_data.pop('info', {}) or {} + resolve_version(model, info) try: ModelInputValidator( model, @@ -129,5 +129,6 @@ django-minio-backend = { git = "https://github.com/theriverman/django-minio-back [tool.ruff.lint.per-file-ignores] "__init__.py" = ["E402", "F401"] +"apps.py" = ["F401"] "**/{tests,docs,tools}/*" = ["E402"] "backend/settings.py" = ["F403", "E402"]