@@ -46,7 +46,7 @@ class GalleryAPIView(APIView): 'images': Image, 'videos': Video, 'audios': Audio, - 'voice_clone': VoiceClone, + 'voice': VoiceClone, }[self.kwargs['strategy']] @extend_schema( @@ -56,7 +56,7 @@ class GalleryAPIView(APIView): str, 'path', required=True, - enum=['images', 'videos', 'audios', 'voice_clone'], + enum=['images', 'videos', 'audios', 'voice'], ), OpenApiParameter('limit', int, required=False), OpenApiParameter('offset', int, required=False), @@ -224,6 +224,39 @@ class ModelVideosAPIView(MediaAPIView): class ModelAudiosAPIVIew(MediaAPIView): manager = Audio + @extend_schema( + parameters=[ + OpenApiParameter('voice_id', UUID, 'query', required=False), + OpenApiParameter('preset_id', UUID, 'query', required=False), + OpenApiParameter('model', str, 'path', required=True), + ], + request=MessageSerializer, + responses={ + 201: MessageSerializer(many=True), + }, + ) + def post(self, request, model: str, *args, **kwargs): + if not (request.FILES.get('file') or request.data.get('file')): + try: + if voice_id := request.query_params.get('voice_id'): + voice = Voice.objects.get(uid=voice_id, user=request.user) + transcription = voice.transcription + elif preset_id := request.query_params.get('preset_id'): + voice = Preset.objects.get(uid=preset_id) + transcription = voice.metadata.get('transcription', '') + else: + voice = Preset.objects.get(slug='russian_1') + transcription = voice.metadata.get('transcription', '') + except (Voice.DoesNotExist, Preset.DoesNotExist): + return Response( + {'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) + class ModelVoiceCloneAPIView(MediaAPIView): manager = VoiceClone @@ -240,22 +273,23 @@ class ModelVoiceCloneAPIView(MediaAPIView): }, ) def post(self, request, model: str, *args, **kwargs): - try: - if voice_id := request.query_params.get('voice_id'): - voice = Voice.objects.get(uid=voice_id, user=request.user) - transcription = voice.transcription - elif preset_id := request.query_params.get('preset_id'): - voice = Preset.objects.get(uid=preset_id) - transcription = voice.metadata.get('transcription', '') - else: - voice = Preset.objects.get(slug='russian_1') - transcription = voice.metadata.get('transcription', '') - except Voice.DoesNotExist: - return Response( - {'detail': _('Voice not found.')}, - status=HTTP_400_BAD_REQUEST, + if not (request.FILES.get('file') or request.data.get('file')): + try: + if voice_id := request.query_params.get('voice_id'): + voice = Voice.objects.get(uid=voice_id, user=request.user) + transcription = voice.transcription + elif preset_id := request.query_params.get('preset_id'): + voice = Preset.objects.get(uid=preset_id) + transcription = voice.metadata.get('transcription', '') + else: + voice = Preset.objects.get(slug='russian_1') + transcription = voice.metadata.get('transcription', '') + except (Voice.DoesNotExist, Preset.DoesNotExist): + return Response( + {'detail': _('Voice not found.')}, + status=HTTP_400_BAD_REQUEST, + ) + request.data.update( + {'file': voice.file, 'info': {'transcription': transcription, **request.data['info']}} ) - request.data.update( - {'file': voice.file, 'info': {'transcription': transcription, **request.data['info']}} - ) return super().post(request, model, *args, **kwargs) @@ -13,5 +13,5 @@ urlpatterns = [ path('image/', ModelImagesAPIView.as_view(), name='images'), path('video/', ModelVideosAPIView.as_view(), name='video'), path('audio/', ModelAudiosAPIVIew.as_view(), name='audio'), - path('voice_clone/', ModelVoiceCloneAPIView.as_view(), name='voice-clone'), + path('voice/', ModelVoiceCloneAPIView.as_view(), name='voice'), ]