@@ -1,8 +1,11 @@ import logging +from typing import Optional +from django.http import QueryDict from django.utils.translation import gettext_lazy as _ from drf_spectacular.utils import OpenApiParameter, extend_schema from rest_framework.permissions import IsAuthenticated +from rest_framework.request import Request from rest_framework.response import Response from rest_framework.status import HTTP_400_BAD_REQUEST, HTTP_402_PAYMENT_REQUIRED from rest_framework.views import APIView @@ -37,6 +40,16 @@ from .models import Audio, Image, Video, VoiceClone, Voice, Preset logger = logging.getLogger(__name__) +def get_request_data(request: Request) -> dict: + data = request.data.copy() + + if isinstance(data, QueryDict): + # Form-data и multipart приходят как QueryDict, а JSON - как обычный словарь. + return data.dict() + + return data + + class GalleryAPIView(APIView): permission_classes = [ IsAuthenticated, @@ -139,9 +152,19 @@ class MediaAPIView(APIView): 200: MessageSerializer(many=True), }, ) - def post(self, request, model: str, *args, **kwargs): + def post( + self, + request: Request, + model: str, + *args, + data: Optional[dict] = None, + **kwargs, + ) -> Response: """Create new media content (image, video, audio) message.""" - serializer = MessageSerializer(data=request.data) + if data is None: + data = request.data + + serializer = MessageSerializer(data=data) if serializer.is_valid(): gallery, created = self.manager.objects.get_or_create( user=request.user, @@ -237,13 +260,17 @@ class ModelAudiosAPIVIew(MediaAPIView): 201: MessageSerializer(many=True), }, ) - def post(self, request, model: str, *args, **kwargs): - if not (request.FILES.get('file') or request.data.get('file')): + def post(self, request: Request, model: str, *args, **kwargs) -> Response: + data = get_request_data(request) + if not (request.FILES.get('file') or data.get('file')): + voice_id = data.pop('voice_id', None) + preset_id = data.pop('preset_id', None) + try: - if voice_id := request.data.pop('voice_id', None): + if voice_id: voice = Voice.objects.get(pk=voice_id, user=request.user) transcription = voice.transcription - elif preset_id := request.data.pop('preset_id', None): + elif preset_id: voice = Preset.objects.get(uid=preset_id) transcription = voice.metadata.get('transcription', '') else: @@ -253,10 +280,10 @@ class ModelAudiosAPIVIew(MediaAPIView): {'detail': _('Voice not found.')}, status=HTTP_400_BAD_REQUEST, ) - request.data.update( - {'file': voice.file, 'info': {'transcription': transcription, **request.data['info']}} - ) - return super().post(request, model, *args, **kwargs) + + info = data['info'] + data.update({'file': voice.file, 'info': {'transcription': transcription, **info}}) + return super().post(request, model, data=data, *args, **kwargs) class ModelVoiceCloneAPIView(MediaAPIView): @@ -271,11 +298,14 @@ class ModelVoiceCloneAPIView(MediaAPIView): 201: MessageSerializer(many=True), }, ) - def post(self, request, model: str, *args, **kwargs): - info = request.data.get('info', {}) or {} - if voice_id := request.data.get('voice_id'): + def post(self, request: Request, model: str, *args, **kwargs) -> Response: + data = get_request_data(request) + info = data.get('info', {}) + + if voice_id := data.get('voice_id'): info.update({'voice_id': voice_id}) - elif preset_id := request.data.get('preset_id'): + elif preset_id := data.get('preset_id'): info.update({'preset_id': preset_id}) - request.data.update({'info': info}) - return super().post(request, model, *args, **kwargs) + data.update({'info': info}) + + return super().post(request, model, data=data, *args, **kwargs)