@@ -84,6 +84,7 @@ TOOLS = [ 'tools.apps.CopywriteConfig', 'tools.apps.MediaConfig', 'tools.apps.FeedConfig', + 'tools.apps.StylesConfig', ] @@ -61,6 +61,7 @@ urlpatterns = [ path('chats/', include('tools.chats.urls')), path('feed/', include('tools.feed.urls')), path('media/', include('tools.media.urls')), + path('styles/', include('tools.styles.urls')), path( 'schema-public/', SpectacularAPIView.as_view(urlconf=['backend.public']), @@ -0,0 +1,3 @@ +""" +This tool provides internal customer interface for using custom LoRA-styles in aggregate with generative models. +""" @@ -0,0 +1,106 @@ +from ninja import Router +from typing import List, Optional +from uuid import UUID + +from django.shortcuts import get_object_or_404 +from django.db.models import Count +from django.db import transaction + +from messages.models import Message + +from .models import Gallery, Conversation +from .tasks import run_multiple_inference +from .schemas import ( + SuccessResponseSchema, + ErrorResponseSchema, + GallerySchema, + CreateConversationSchema, + ConversationSchema, +) + +router = Router() + + +@router.post('/gallery/init/', response={200: GallerySchema, 201: GallerySchema, 400: ErrorResponseSchema}) +def init_user_gallery(request): + """Единожды создает хранилище, если еще не создано.""" + try: + with transaction.atomic(): + gallery, created = Gallery.objects.get_or_create(user=request.user) + + if created: + return 201, gallery + return 200, gallery + + except Exception as e: + return 400, ErrorResponseSchema(error='Ошибка инициализации хранилища.', details={'traceback': e}) + + +@router.get( + '/gallery/', + response={200: GallerySchema, 404: ErrorResponseSchema}, + summary='Информация о хранилище пользователя.', +) +def get_user_gallery(request): + """Возвращает хранилище текущего пользователя.""" + gallery = get_object_or_404(Gallery, user=request.user) + return 200, gallery + + +@router.post( + '/conversations/', + response={201: ConversationSchema, 400: ErrorResponseSchema}, + summary='Создание нового набора для обработки.', +) +def create_conversation(request, data: CreateConversationSchema): + """Создает новый набор про входящему промпту и обрабатывает его стилем.""" + gallery = get_object_or_404(Gallery, user=request.user) + + conversation = Conversation.objects.create(gallery=gallery, style_slug=data.style_slug) + + prompt_message = get_object_or_404( + Message, uid=data.prompt_uid, content_object__gallery__user=request.user + ) + + try: + run_multiple_inference.delay( + user_id=request.user.id, + conversation_uid=conversation.uid, + inference_slug=data.style_slug, # нужна проверка заранее? + input_message_id=data.prompt_uid, + count=data.images_count, + ) + + return 201, conversation + + except Exception as e: + conversation.delete() + return 400, ErrorResponseSchema(error='Ошибка генерации изображений', details={'traceback': str(e)}) + + +@router.get( + '/conversations/ready/', response=List[ConversationSchema], summary='Список наборов, готовых к обработке' +) +def get_ready_conversations(request, style_slug: Optional[str] = None): + """ + Возвращает наборы сообщениями, готовые к обработке. + """ + queryset = Conversation.objects.annotate(messages_count=Count('messages')).filter( + messages_count__gte=2, messages_count__lte=4 + ) + + if style_slug: + queryset = queryset.filter(style_slug=style_slug) + + return queryset.order_by('-created_at') + + +@router.get( + '/conversations/{conversation_uid}/', + response={200: ConversationSchema, 404: ErrorResponseSchema}, + summary='Информация о конкретном наборе', +) +def get_conversation_details(request, conversation_uid: UUID): + """Возвращает детальную информацию о наборе""" + conversation = get_object_or_404(Conversation, uid=conversation_uid, gallery__user=request.user) + return 200, conversation @@ -0,0 +1,70 @@ +from uuid import uuid4 + +from django.db import models +from django.core.exceptions import ValidationError +from django.contrib.contenttypes.fields import GenericRelation +from django.db.models.query import QuerySet + +from messages.models import SingleStore, Message + + +class Gallery(SingleStore): + class Meta: + verbose_name = 'Хранилище' + verbose_name_plural = 'Хранилища' + + def __str__(self): + return f'Хранилище пользователя {self.user_id}' + + @property + def output_messages(self) -> QuerySet[Message]: + return self.messages.filter(from_model=True) + + +class Conversation(models.Model): + """ + Pack of some (2-4) images with LoRA-styles application. + """ + + uid = models.UUIDField(primary_key=True, default=uuid4, editable=False, verbose_name='Идентификатор') + + gallery = models.ForeignKey( + Gallery, on_delete=models.CASCADE, related_name='conversations', verbose_name='Хранилище' + ) + + style_slug = models.CharField(max_length=50) + + messages = GenericRelation( + Message, + content_type_field='content_type', + object_id_field='object_id', + related_query_name='conversation', + verbose_name='Сообщения', + ) + + created_at = models.DateTimeField(auto_now_add=True, verbose_name='Дата создания') + updated_at = models.DateTimeField(auto_now=True, verbose_name='Дата обновления') + + class Meta: + verbose_name = 'Набор стилей' + verbose_name_plural = 'Наборы стилей' + ordering = ['-created_at'] + + @property + def messages_count(self) -> int: + return self.messages.count() + + @property + def is_ready_for_processing(self) -> bool: + return self.messages.count() <= 10 + + def __str__(self): + return f'Набор #{self.uid} в хранилище {self.gallery.uid}' + + def clean(self): + count = self.messages.count() + if not (2 <= count <= 4): + raise ValidationError('Набор должен содержать до 10 изображений!') + + def save(self, *args, **kwargs): + super().save(*args, **kwargs) @@ -0,0 +1,46 @@ +from ninja import Schema, ModelSchema +from uuid import UUID +from typing import Optional, List +from .models import Gallery, Conversation + + +# Base Schemas +class SuccessResponseSchema(Schema): + success: bool = True + message: str = 'Операция выполнена успешно' + data: Optional[dict] = None + + +class ErrorResponseSchema(Schema): + success: bool = False + error: str + details: Optional[dict] = None + + +# Gallery Schemas +class GallerySchema(ModelSchema): + active_conversations: int + + class Config: + model = Gallery + fields = ['uid', 'created_at'] + + @staticmethod + def resolve_active_conversations(obj) -> int: + return obj.conversations.count() + + +# Conversation Schemas +class CreateConversationSchema(Schema): + style_slug: str + prompt_uid: UUID + images_count: int + + +class ConversationSchema(ModelSchema): + messages_count: int + is_ready_for_processing: bool + + class Config: + model = Conversation + fields = ['uid', 'style_slug', 'created_at', 'updated_at'] @@ -0,0 +1,46 @@ +from celery import shared_task +from uuid import UUID + +from django.db import transaction +from django.core.cache import cache + +from messages.models import Message +from .models import Conversation +from ml_model.tasks import _run_inference + + +@shared_task +def run_multiple_inference( + user_id: UUID, + conversation_uid: UUID, + inference_slug: str, + input_message_id: UUID, + count: int, + history_ids: list[UUID] = list(Message.objects.none().values('pk')), +): + input_message = Message.objects.get(uid=input_message_id) + conversation = Conversation.objects.get(uid=conversation_uid) + + cache_key = f'{input_message.content_object._meta.model_name}s:{input_message.content_object.uid}' + cache.add(cache_key, []) + + with transaction.atomic(): + output_slots = [] + for _ in range(count): + slot = Message.objects.create( + from_model=True, + content_object=conversation, + info={'style_slug': inference_slug, 'conversation_uid': str(conversation_uid)}, + ) + output_slots.append(slot.uid) + + for slot_uid in output_slots: + _run_inference.delay( + user_id=user_id, + inference_slug=inference_slug, + input_message_id=input_message_id, + output_slot_id=slot_uid, + history_ids=history_ids, + ) + + return output_slots @@ -0,0 +1,15 @@ +from ninja import NinjaAPI +from django.urls import path + +from .apis import router + +api = NinjaAPI( + title='Styles Module', + version='1.0.0', +) + +api.add_router('/styles/', router) + +urlpatterns = [ + path('', api.urls), +] @@ -36,3 +36,9 @@ class FeedConfig(AppConfig): default_auto_field = 'django.db.models.BigAutoField' name = 'tools.feed' verbose_name = _('Feed') + + +class StylesConfig(AppConfig): + default_auto_field = 'django.db.models.BigAutoField' + name = 'tools.styles' + verbose_name = _('Styles')