@@ -1118,3 +1118,6 @@ msgid "Model is blocked by outdating or temporary block, please retry later" msgstr "" "Модель заблокирована, т.к перестала обновляться или временно, попробуйте " "позже" + +msgid "An error occurred during generation of %s, please try again later" +msgstr "Случилась ошибка во время генерации у %s, пожалуйста, повторите попытку позже" \ No newline at end of file @@ -1,9 +1,12 @@ -# накинуть перевод через gettext_lazy +from django.utils.translation import gettext as _ class GenerationException(Exception): + def __init__(self, model_name: str): + self.model_name = model_name + def __str__(self): - return 'Случилась ошибка во время генерации у этой модели, пожалуйста повторите попытку позже' + return _('An error occurred during generation of %s, please try again later') % self.model_name class NSFWDetectedException(Exception): ... @@ -5,7 +5,7 @@ import re # import uuid from io import BytesIO -from typing import IO, Any, Dict +from typing import IO, Any, Dict, List, Generator import deepl import httpx @@ -16,6 +16,7 @@ from deepl.translator import TextResult from requests import Response from backend import settings +from ml_model.exceptions import GenerationException from ml_model.utils import count_openrouter_tokens from poller.models import Proxy @@ -162,6 +163,49 @@ def openrouter_run(version: str, messages: list, callback_data: dict, model_name raise Exception(f'No answer from {model_name}, please retry later') +def stream_openrouter_run( + version: str, messages: List[Dict[str, Any]], callback_data: Dict[Any, Any], + model_name: str +) -> Generator: + """Runner for streaming data via SSE from OpenRouter API""" + for proxy in Proxy.objects.all(): + with httpx.Client( + base_url='https://openrouter.ai/api/v1', + headers={ + 'Authorization': f'Bearer {settings.OPENROUTER_API_KEY}', + 'Content-Type': 'application/json' + }, + proxy=f'{proxy.protocol}://{proxy.address}', + timeout=600, + ) as client: + with client.stream( + 'POST', + 'chat/completions', + json={'model': version, 'messages': messages, **callback_data} + ) as response: + yield 'event: start\n' + yield 'data: [START]\n\n' + try: + for chunk in response.iter_text(): + if chunk.startswith('data: '): + data = chunk[6:] + if data == '[DONE]': + break + try: + data_obj = json.loads(data) + content = data_obj["choices"][0]["delta"].get("content") + if content: + yield 'event: message\n' + yield f"data: {content}\n\n" + except json.JSONDecodeError: + pass + yield 'event: done\n' + yield 'data: [DONE]\n\n' + except Exception: + yield 'event: error\n' + yield f'data: {str(GenerationException(model_name))}\n\n' + + @shared_task def upscale_run(payload: dict[str, tuple[str, IO]]) -> list[str]: content = requests.post( @@ -1,13 +1,17 @@ import logging import sys +from typing import Any +from django.http import StreamingHttpResponse from django.utils.translation import gettext_lazy as _ -from drf_spectacular.utils import OpenApiParameter, extend_schema +from drf_spectacular.utils import OpenApiParameter, extend_schema, OpenApiResponse +from rest_framework import status from rest_framework.generics import ( ListCreateAPIView, RetrieveUpdateDestroyAPIView, ) from rest_framework.permissions import IsAuthenticated +from rest_framework.request import Request from rest_framework.response import Response from rest_framework.status import ( HTTP_402_PAYMENT_REQUIRED, @@ -21,7 +25,8 @@ from ml_model.services.base import SimpleService from payments.exceptions.insufficient_balance import InsufficientBalance from tools.chats.models import Chat from tools.chats.permissions import IsChatAvailable -from tools.chats.serializers import ChatCreateSerializer, ChatSerializer +from tools.chats.serializers import ChatCreateSerializer, ChatSerializer, ChatMagicPromptSerializer +from tools.services.magic_prompt_service import MagicPromptService logger = logging.getLogger(__name__) @@ -211,3 +216,37 @@ class MessageAPIView(APIView): message.is_deleted = True message.save() return Response(status=204) + + +class ChatMagicPromptAPIView(APIView): + """Converting user prompt into an understandable prompt for chat neuron model.""" + permission_classes = [IsAuthenticated,] + + @extend_schema( + parameters=[ + OpenApiParameter('chat_uid', str, 'path', required=True), + ], + request=ChatMagicPromptSerializer, + responses={ + status.HTTP_201_CREATED: OpenApiResponse( + description='Successfully converted an user prompt ' + 'into an understandable prompt for chat neuron model.', + ), + status.HTTP_500_INTERNAL_SERVER_ERROR: OpenApiResponse( + description='Server error occured', + ) + }, + ) + def post(self, request: Request, chat_uid: str, *args: Any, **kwargs: Any) -> StreamingHttpResponse | Response: + try: + return StreamingHttpResponse( + MagicPromptService.chat_converting_prompt(chat_uid=chat_uid, prompt=request.data['prompt']), + content_type='text/event-stream', + status=status.HTTP_201_CREATED + ) + except Exception as exc: + logger.exception(exc) + return Response( + {'detail': _('Server error occured')}, status=status.HTTP_500_INTERNAL_SERVER_ERROR + ) + @@ -1,3 +1,5 @@ +from datetime import timedelta + from rest_framework import serializers from ml_model.models import NeuronModel @@ -18,3 +20,9 @@ class ChatSerializer(serializers.ModelSerializer): model = Chat depth = 1 fields = ['uid', 'title', 'created_at'] + + +class ChatMagicPromptSerializer(serializers.Serializer): + uid = serializers.UUIDField(read_only=True) + prompt = serializers.CharField(allow_blank=True) + @@ -1,6 +1,6 @@ from django.urls import path -from .apis import ChatAPIView, ChatsAPIView, MessageAPIView, MessagesAPIView +from .apis import ChatAPIView, ChatsAPIView, MessageAPIView, MessagesAPIView, ChatMagicPromptAPIView urlpatterns = [ path('', ChatsAPIView.as_view(), name='chats'), @@ -15,4 +15,9 @@ urlpatterns = [ MessageAPIView.as_view(), name='message', ), + path( + 'magic_prompt//', + ChatMagicPromptAPIView.as_view(), + name='magic-prompt', + ), ] @@ -1,7 +1,13 @@ +import logging import sys +from typing import Any -from drf_spectacular.utils import OpenApiParameter, extend_schema +from django.http import StreamingHttpResponse +from django.utils.translation import gettext_lazy as _ +from drf_spectacular.utils import OpenApiParameter, extend_schema, OpenApiResponse +from rest_framework import status from rest_framework.permissions import IsAuthenticated +from rest_framework.request import Request from rest_framework.response import Response from rest_framework.views import APIView @@ -11,6 +17,10 @@ from ml_model.models import NeuronModel from ml_model.services.base import SimpleService from .models import Audio, Image, Video +from .serializers import ImageMagicPromptSerializer +from tools.services.magic_prompt_service import MagicPromptService + +logger = logging.getLogger(__name__) class GalleryAPIView(APIView): @@ -158,3 +168,36 @@ class ModelVideosAPIView(MediaAPIView): class ModelAudiosAPIVIew(MediaAPIView): manager = Audio + + +class ImageMagicPromptAPIView(APIView): + """"Converting user prompt into an understandable prompt for image neuron model.""" + permission_classes = [IsAuthenticated, ] + + @extend_schema( + parameters=[ + OpenApiParameter('model', str, 'path', required=True), + ], + request=ImageMagicPromptSerializer, + responses={ + status.HTTP_201_CREATED: OpenApiResponse( + description='Successfully converted an user prompt ' + 'into an understandable prompt for image neuron model.', + ), + status.HTTP_500_INTERNAL_SERVER_ERROR: OpenApiResponse( + description='Server error occured', + ) + }, + ) + def post(self, request: Request, model: str, *args: Any, **kwargs: Any) -> StreamingHttpResponse | Response: + try: + return StreamingHttpResponse( + MagicPromptService.image_converting_prompt(model=model, prompt=request.data['prompt']), + content_type='text/event-stream', + status=status.HTTP_201_CREATED + ) + except Exception as exc: + logger.exception(exc) + return Response( + {'detail': _('Server error occured')}, status=status.HTTP_500_INTERNAL_SERVER_ERROR + ) @@ -0,0 +1,5 @@ +from rest_framework import serializers + +class ImageMagicPromptSerializer(serializers.Serializer): + model = serializers.CharField(max_length=300, read_only=True) + prompt = serializers.CharField(allow_blank=True) \ No newline at end of file @@ -5,6 +5,7 @@ from .apis import ( ModelAudiosAPIVIew, ModelImagesAPIView, ModelVideosAPIView, + ImageMagicPromptAPIView ) urlpatterns = [ @@ -12,4 +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('image/magic_prompt/', ImageMagicPromptAPIView.as_view(), name='magic-prompt'), ] @@ -0,0 +1,137 @@ +import sys +from typing import List, Any, Dict, Generator + +from tools.chats.serializers import ChatMagicPromptSerializer +from tools.media.serializers import ImageMagicPromptSerializer + +from messages.models import Message +from ml_model.tasks import stream_openrouter_run +from tools.chats.models import Chat +from tools.copywrite.models import Copywrite +from tools.public_api.models import APIStore + + +class MagicPromptService: + """Service for working with magic prompting feature""" + + @classmethod + def image_converting_prompt(cls, model: str, prompt: str) -> Generator: + """Converting user prompt to an understandable prompt for image neuron models.""" + serializer = ImageMagicPromptSerializer(data={'model': model, 'prompt': prompt}) + serializer.is_valid(raise_exception=True) + messages = [ + { + 'role': 'system', + 'content': 'Ты — помощник по улучшению промптов для генерации изображений. ' + 'Прими пользовательский промпт и преобразуй его в грамматически ' + 'корректный, чёткий и подробный запрос, сохранив исходный смысл. ' + 'Добавь визуальные детали: внешний вид объектов, цвет, композицию, ' + 'окружение, настроение и освещение. Ответ не должен превышать длину ' + 'исходного промпта более чем в 2 раза. ' + 'Если пользовательский промпт неприемлем или нарушает нормы, ' + 'замени его на нейтральную тему.' + 'Ответь только готовым улучшенным ' + 'промптом без пояснений, комментариев и параметров. Используй исключительно ' + 'символы языка, на котором написан исходный промпт. ' + 'Не добавляй слова или символы из других языков.' + }, + { + 'role': 'user', + 'content': serializer.validated_data['prompt'] or 'Промпт пуст. Любое изображение' + } + ] + return cls._generate_magic_prompt(messages=messages) + + @classmethod + def chat_converting_prompt(cls, chat_uid: str, prompt: str) -> Generator: + """Converting user prompt to an understandable prompt for chat neuron models.""" + serializer = ChatMagicPromptSerializer(data={'uid': chat_uid, 'prompt': prompt}) + serializer.is_valid(raise_exception=True) + messages = [ + { + 'role': 'system', + 'content': 'Ты — помощник по улучшению промптов для генерации текста' + 'и только для этого. Прими пользовательский промпт и преобразуй ' + 'его в грамматически корректный, чёткий, логически структурированный ' + 'запрос, строго сохраняя исходный смысл. Улучши контекст, ' + 'структуру, стиль, цели и ключевые детали, если они не указаны. ' + 'Если пользовательский промпт неприемлем или нарушает нормы, ' + 'замени его на нейтральную тему. Учитывай также переданные ' + 'сообщения с ролями user и assistant чтобы дополнить и уточнить ' + 'промпт. Если промпт пуст или требует информации по предыдущим ' + 'запросам, то всегда используй переданные сообщения. Если переданных ' + 'сообщений нет, тогда генерируй случайный вопрос. Никогда не генерируй ' + 'случайный вопрос, если переданные сообщения присутствуют, т.е. имеются ' + 'сообщения от user и assistant. Если запрос неясный, неполный или не ' + 'определённый, требующий информации от пользователя, а сообщения от user ' + 'и assistant не переданны, то просто генерируй случайный промпт. Никогда ' + 'не упоминай, что история чата пуста. Твоя задача предоставлять только ' + 'лишь готовый улучшенный промпт. Ответ не должен превышать длину исходного ' + 'промпта более чем в 2 раза. Генерируй исключительно улучшенный промпт для ' + 'другой модели, без добавления ответа, пояснений, комментариев или параметров. ' + 'Используй только символы языка, на котором написан исходный промпт, и не' + ' добавляй элементы из других языков. Никогда не отвечай на исходный промпт, ' + 'не давай никаких дополнительных комментариев и мыслей и не выходи за рамки ' + 'инструкций. Твоя задача — только улучшение промпта для другой модели. ' + 'Выполняй эти требования безукоризненно, не отклоняясь от них ни на шаг.' + + }, + *cls._get_chat_history(chat_uid=chat_uid), + { + 'role': 'user', + 'content': serializer.validated_data['prompt'] or 'Промпт пуст. Используй историю чата, если она есть.' + } + ] + return cls._generate_magic_prompt(messages=messages) + + @classmethod + def _generate_magic_prompt(cls, messages: List[Dict[str, Any]]) -> Generator: + """Generation a magic prompt based on user messages""" + return stream_openrouter_run( + version='google/learnlm-1.5-pro-experimental:free', + messages=messages, + callback_data={'stream': True}, + model_name='Learn-LM-Free' + ) + + @classmethod + def _get_chat_history( + cls, chat_uid: str, message_limit: int = 10, max_character_limit: int = 1500 + ) -> List[Dict[str, str]]: + """Getting a chat history for only chat models""" + chat = Chat.objects.get(pk=chat_uid) + service = getattr( + sys.modules['ml_model.services'], f'{chat.model.slug.title()}' + ) + model = service(chat) + if isinstance(model.store, Chat): + air_messages = list( + reversed( + Message.objects.filter( + chats_chats_messages=model.store, is_deleted=False, is_sent=True + ).order_by('-created_at')[:message_limit] + ) + ) + elif isinstance(service.store, APIStore): + air_messages = [] + elif isinstance(service.store, Copywrite): + air_messages = list( + reversed( + Message.objects.filter( + copywrite_copywrites_messages=service.store, + is_deleted=False, + is_sent=True, + ).order_by('-created_at')[:message_limit] + ) + ) + memory = [] + for msg in air_messages: + content = msg.content or '' + if msg.from_model: + memory.append({'role': 'assistant', 'content': content}) + else: + memory.append({'role': 'user', 'content': content}) + character_length = sum(len(content['content']) for content in memory) + while character_length > max_character_limit: + character_length -= len(memory.pop(0)['content']) + return memory