@@ -6,7 +6,9 @@ from rest_framework.views import APIView from messages.models import Message from messages.serializers import MessageSerializer +from ml_model.exceptions import FileNotProvided, InvalidParameterError from ml_model.models import ModelParameter +from ml_model.validators import ModelInputValidator from tools.chats.models import Chat @@ -34,7 +36,7 @@ class MessagesAPIView(APIView): """Create Message with ml_model in chat""" serializer = MessageSerializer(data=request.data) if serializer.is_valid(): - chat = Chat.objects.get(pk=chat_uid) + chat = Chat.objects.select_related('model').get(pk=chat_uid) info = {} if chat.model: service = chat.model.service @@ -54,6 +56,16 @@ class MessagesAPIView(APIView): case 'list': missing_info.update({p.key: p.default.split(',') if p.default else []}) merged_info = info | missing_info + try: + ModelInputValidator( + chat.model, + content=serializer.validated_data.get('content'), + file=serializer.validated_data.get('file'), + info=merged_info, + ).validate() + except (FileNotProvided, InvalidParameterError) as exc: + return Response({'detail': str(exc)}, status=400) + i = Message.objects.create( **serializer.validated_data, info=merged_info, @@ -1,5 +1,6 @@ from typing import Iterable +from django.utils.functional import Promise from django.utils.translation import gettext as _ # накинуть перевод через gettext_lazy @@ -140,11 +141,11 @@ class PredictionInterruptedError(Exception): class InvalidParameterError(Exception): - def __init__(self, error_text: str): + def __init__(self, error_text: str | Promise): self.error_text = error_text - def __str__(self): - return self.error_text + def __str__(self) -> str: + return str(self.error_text) class PromptLengthExceeded(Exception): def __init__(self, max_length: int = 3000) -> None: @@ -0,0 +1,58 @@ +from typing import Any + +from django.db.models import Q +from django.utils.translation import gettext_lazy as _ + +from ml_model.exceptions import FileNotProvided, InvalidParameterError +from ml_model.models import ModelInput, NeuronModel + + +class ModelInputValidator: + def __init__( + self, + model: NeuronModel, + *, + content: str | None, + file: Any, + info: dict | None = None, + ) -> None: + self.model = model + self.content = content + self.file = file + self.info = info or {} + + def validate(self) -> None: + required_input_types = self._get_required_input_types() + + if ModelInput.TypeChoices.TEXT.value in required_input_types: + self._validate_text_input() + + required_file_input_types = required_input_types - {ModelInput.TypeChoices.TEXT.value} + if required_file_input_types: + self._validate_file_input(required_file_input_types) + + def _get_required_input_types(self) -> set[str]: + version = self.info.get('version') + required_inputs = self.model.inputs.filter(required=True) + if version: + required_inputs = required_inputs.filter( + Q(versions__isnull=True) | Q(versions__slug=version) + ).distinct() + else: + required_inputs = required_inputs.filter(versions__isnull=True) + + return set(required_inputs.values_list('type', flat=True)) + + def _validate_text_input(self) -> None: + if not self._has_text_content(): + raise InvalidParameterError(_('The request must not be empty')) + + def _validate_file_input(self, required_file_input_types: set[str]) -> None: + if not self.file: + input_type = sorted(required_file_input_types)[0] + input_label = ModelInput.TypeChoices(input_type).label + + raise FileNotProvided(input_label) + + def _has_text_content(self) -> bool: + return isinstance(self.content, str) and bool(self.content.strip()) @@ -8,8 +8,10 @@ from ninja.errors import HttpError from authentication.security import SyncAuthBearer from messages.models import Message +from ml_model.exceptions import FileNotProvided, InvalidParameterError from ml_model.models import NeuronModel from ml_model.schemas import NeuronModelLink +from ml_model.validators import ModelInputValidator from tools.chats.models import Chat from tools.chats.schemas import MessageInSchema from tools.chats.services.sse_chat_stream import SSEChatStreamService @@ -58,6 +60,16 @@ def stream_message(request, chat_uid: UUID, body: MessageInSchema): raise HttpError(409, _('Stream already in progress')) data = body.model_dump(include={'content', 'file', 'info'}, exclude_unset=True) + try: + ModelInputValidator( + chat.model, + content=data.get('content'), + file=data.get('file'), + info=data.get('info'), + ).validate() + except (FileNotProvided, InvalidParameterError) as exc: + raise HttpError(400, str(exc)) + input_message = Message.objects.create(content_object=chat, from_model=False, **data) store.start() @@ -23,6 +23,7 @@ from ml_model.exceptions import ( DeploymentDisabled, ExceededContextLengthError, FileExtensionNotSupported, + FileNotProvided, FileTooLargeError, FileUploadUnsupported, ImageAnalysisError, @@ -36,6 +37,7 @@ from ml_model.exceptions import ( TemplateUnknownException, UnrecognizedFileError, ) +from ml_model.validators import ModelInputValidator from payments.exceptions.insufficient_balance import InsufficientBalance from tools.chats.models import Chat from tools.chats.permissions import IsChatAvailable @@ -153,9 +155,19 @@ class MessagesAPIView(APIView): """ serializer = MessageSerializer(data=request.data) if serializer.is_valid(): - chat = Chat.objects.get(pk=chat_uid) + chat = Chat.objects.select_related('model').get(pk=chat_uid) info = serializer.validated_data.pop('info', {}) service = chat.model.service + try: + ModelInputValidator( + chat.model, + content=serializer.validated_data.get('content'), + file=serializer.validated_data.get('file'), + info=info, + ).validate() + except (FileNotProvided, InvalidParameterError) as exc: + return Response({'detail': str(exc)}, status=HTTP_400_BAD_REQUEST) + input_message = Message.objects.create( **serializer.validated_data, info=info, @@ -166,11 +178,11 @@ class MessagesAPIView(APIView): output_messages = service(chat).make(input_message) except DeploymentDisabled as exc: return Response( - {'detail': f'{exc}'}, + {'detail': str(exc)}, status=HTTP_503_SERVICE_UNAVAILABLE, ) except PaidPlanRequiredError as exc: - return Response({'detail': f'{exc}'}, status=HTTP_402_PAYMENT_REQUIRED) + return Response({'detail': str(exc)}, status=HTTP_402_PAYMENT_REQUIRED) except ( FileExtensionNotSupported, ExceededContextLengthError, @@ -185,17 +197,17 @@ class MessagesAPIView(APIView): ModelVersionNotAvailable, OutputSensitiveImageContentError, ) as exc: - return Response({'detail': f'{exc}'}, status=HTTP_400_BAD_REQUEST) + return Response({'detail': str(exc)}, status=HTTP_400_BAD_REQUEST) except TemplateNotFound as exc: - return Response({'detail': f'{exc}'}, status=HTTP_500_INTERNAL_SERVER_ERROR) + return Response({'detail': str(exc)}, status=HTTP_500_INTERNAL_SERVER_ERROR) except TemplateUnknownException as exc: logger.exception(exc) - return Response({'detail': f'{exc}'}, status=HTTP_500_INTERNAL_SERVER_ERROR) + return Response({'detail': str(exc)}, status=HTTP_500_INTERNAL_SERVER_ERROR) except Exception as exc: input_message.is_sent = False input_message.save() if isinstance(exc, InsufficientBalance): - return Response({'detail': f'{exc}'}, status=HTTP_402_PAYMENT_REQUIRED) + return Response({'detail': str(exc)}, status=HTTP_402_PAYMENT_REQUIRED) logger.exception(exc) return Response( { @@ -35,6 +35,7 @@ from ml_model.exceptions import ( UnsupportedSize, ) from ml_model.models import NeuronModel +from ml_model.validators import ModelInputValidator from payments.exceptions.insufficient_balance import InsufficientBalance from .models import Audio, Image, Video, VoiceClone, Voice, Preset @@ -126,7 +127,7 @@ class MediaAPIView(APIView): """ List Messages """ - gallery, _ = self.manager.objects.get_or_create( + gallery, _ = self.manager.objects.select_related('model').get_or_create( user=request.user, model__slug=model, defaults={ @@ -170,7 +171,7 @@ class MediaAPIView(APIView): serializer = MessageSerializer(data=data) if serializer.is_valid(): - gallery, created = self.manager.objects.get_or_create( + gallery, created = self.manager.objects.select_related('model').get_or_create( user=request.user, model__slug=model, defaults={ @@ -180,6 +181,16 @@ class MediaAPIView(APIView): ) info = serializer.validated_data.pop('info', {}) service = gallery.model.service + try: + ModelInputValidator( + gallery.model, + content=serializer.validated_data.get('content'), + file=serializer.validated_data.get('file'), + info=info, + ).validate() + except (FileNotProvided, InvalidParameterError) as exc: + return Response({'detail': str(exc)}, status=HTTP_400_BAD_REQUEST) + input_message = Message.objects.create( **serializer.validated_data, info=info, @@ -209,12 +220,12 @@ class MediaAPIView(APIView): RealPersonDetectedError, OutputSensitiveImageContentError, ) as exc: - return Response({'detail': f'{exc}'}, status=HTTP_400_BAD_REQUEST) + return Response({'detail': str(exc)}, status=HTTP_400_BAD_REQUEST) except Exception as exc: input_message.is_sent = False input_message.save() if isinstance(exc, InsufficientBalance): - return Response({'detail': f'{exc}'}, status=HTTP_402_PAYMENT_REQUIRED) + return Response({'detail': str(exc)}, status=HTTP_402_PAYMENT_REQUIRED) logger.exception(exc) if any(phrase in str(exc) for phrase in ('Insufficient credit', 'Request was throttled')): return Response( @@ -7,7 +7,9 @@ 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.validators import ModelInputValidator from tools.chats.schemas import MessageInSchema from tools.chats.services.sse_chat_stream import SSEChatStreamService @@ -105,10 +107,17 @@ def public_stream_message(request, model_slug: str, body: MessageInSchema): if not model.streaming: raise HttpError(501, _('Stream not supported for this model')) - if not body.content: - raise HttpError(400, _('The request must not be empty')) - data = body.model_dump(include={'content', 'file', 'info'}, exclude_unset=True) + try: + ModelInputValidator( + model, + content=data.get('content'), + file=data.get('file'), + info=data.get('info'), + ).validate() + except (FileNotProvided, InvalidParameterError) as exc: + raise HttpError(400, str(exc)) + input_message = Message.objects.create( content_object=api_store, from_model=False, from_public_api=True, **data ) @@ -14,10 +14,11 @@ from rest_framework.views import APIView from messages.models import Message from messages.serializers import MessageSerializer from ml_model.choices import ContentTypes -from ml_model.exceptions import InvalidParameterError +from ml_model.exceptions import FileNotProvided, InvalidParameterError from ml_model.models import NeuronModel from ml_model.selectors.ml_models_selector import NeuronModelSelector from ml_model.serializers import PublicNeuronModelSerializer +from ml_model.validators import ModelInputValidator from payments.exceptions.insufficient_balance import InsufficientBalance from tools.public_api.models import APIKey, APIStore from tools.public_api.permissions import HasAPIKey @@ -70,13 +71,18 @@ class BaseGenerationView(APIView): ) serializer = MessageSerializer(data=request.data) serializer.is_valid(raise_exception=True) - if not serializer.validated_data.get('content'): - return Response( - {'detail': _('The request must not be empty')}, - status=HTTP_400_BAD_REQUEST, - ) service = model.service info = serializer.validated_data.pop('info', {}) + try: + ModelInputValidator( + model, + content=serializer.validated_data.get('content'), + file=serializer.validated_data.get('file'), + info=info, + ).validate() + except (FileNotProvided, InvalidParameterError) as exc: + return Response({'detail': str(exc)}, status=HTTP_400_BAD_REQUEST) + input_message = Message.objects.create( **serializer.validated_data, info=info, @@ -92,9 +98,9 @@ class BaseGenerationView(APIView): input_message.is_sent = False input_message.save() if isinstance(exc, InsufficientBalance): - return Response({'detail': f'{exc}'}, status=HTTP_402_PAYMENT_REQUIRED) + return Response({'detail': str(exc)}, status=HTTP_402_PAYMENT_REQUIRED) elif isinstance(exc, (InvalidParameterError, ValidationError)): - return Response({'detail': f'{exc}'}, status=HTTP_400_BAD_REQUEST) + return Response({'detail': str(exc)}, status=HTTP_400_BAD_REQUEST) logger.exception(exc) return Response( { @@ -27,3 +27,7 @@ celerybeat-schedule .vscode/ scripts/ + +# Codex local instructions +AGENTS.md +agents.md