@@ -257,6 +257,9 @@ SPECTACULAR_SETTINGS = { # Redis REDIS_HOST = env.str('REDIS_HOST', 'cache-mdb') REDIS_PORT = env.int('REDIS_PORT', 6379) +SSE_STREAM_TTL = env.int('SSE_STREAM_TTL', 600) +SSE_DONE_STREAM_TTL = env.int('SSE_DONE_STREAM_TTL', 30) +SSE_POLL_INTERVAL = env.float('SSE_POLL_INTERVAL', 0.2) # Celery CELERY_BROKER_URL = env.str('CELERY_BROKER_URL', 'redis://celery-mdb:6379/0') @@ -15,12 +15,12 @@ from authentication.exceptions import ( InvalidToken, ) from backend.public import urlpatterns as public_urlpatterns - +from lib.parsers import MultiContentTypeParser logger = logging.getLogger(__name__) -api = NinjaAPI(title='AIR API', version='1.0.0', docs_url=None) +api = NinjaAPI(title='AIR API', version='1.0.0', parser=MultiContentTypeParser(), docs_url=None) compatibility_api = NinjaAPI(title='AIR API DEBUG', version='0.0.1', docs_url=None) compatibility_api_v2 = NinjaAPI(title='AIR API DEBUG v2', version='2.0.0', docs_url=None) @@ -0,0 +1,22 @@ +import orjson + +from django.utils.translation import gettext as _ + +from ninja.errors import HttpError +from ninja.parser import Parser + + +class MultiContentTypeParser(Parser): + def parse_body(self, request): + ct = request.content_type.lower() + if ct.startswith('application/json'): + if not request.body: + return {} + try: + return orjson.loads(request.body) + except orjson.JSONDecodeError: + raise HttpError(400, _('Invalid JSON payload')) + elif ct.startswith('multipart/form-data'): + data = request.POST.dict() + data.update(request.FILES.dict()) + return data @@ -1509,6 +1509,14 @@ msgstr "Голоса" msgid "Voice not found" msgstr "Голос не найден" +#: tools/chats/routes/v1.py:36 +msgid "Chat not found" +msgstr "Чат не найден" + +#: tools/chats/routes/v1.py:41 +msgid "Stream not supported for this model" +msgstr "Стриминг не поддерживается для этой модели" + #: tools/public_api/exceptions.py:7 msgid "Upgrade token limit on your api-key" msgstr "Необходимо повысить лимит токенов у API-ключа" @@ -74,6 +74,7 @@ from ml_model.services.sora import Sora from ml_model.services.stablediffusion import Stablediffusion from ml_model.services.stablemusic import Stablemusic from ml_model.services.suno import Suno +from ml_model.services.text_test_model import Text_Test_Model from ml_model.services.upscaleai import Upscaleai from ml_model.services.veo import Veo from ml_model.services.vicuna import Vicuna @@ -1,6 +1,6 @@ from abc import ABC, abstractmethod from decimal import Decimal -from typing import Never, Any +from typing import Any, Generator, Never from asgiref.sync import async_to_sync from googletrans import Translator @@ -82,3 +82,8 @@ class SimpleService(ABC): @classmethod def predict_price(cls, content: str, file_exists: bool, info: dict[str, Any]) -> Decimal | None: return None + + +class StreamSimpleService(SimpleService): + @abstractmethod + def make_stream(self, input_message: Message, save: bool = True) -> Generator: ... @@ -0,0 +1,128 @@ +import hashlib + +import time +import tiktoken + +from dataclasses import dataclass +from datetime import timedelta +from decimal import Decimal +from django.core.cache import cache + +from messages.models import Message +from ml_model.services.base import StreamSimpleService + +from typing import Iterator + +type TestTextTokens = list[str] + + +@dataclass +class PrepareTestData: + input_tokens: TestTextTokens + output_tokens: TestTextTokens + ttft: float + tbt: float + start_time: float + + +class Text_Test_Model(StreamSimpleService): + TOKENS_COST = { + 'input': Decimal('1000'), # per 1 million input-tokens + 'output': Decimal('2500'), # per 1 million output-tokens + } + + ENCODING = 'o200k_base' + + BASE_OUTPUT_MESSAGE = """ + Lorem ipsum dolor sit amet, consectetur adipiscing elit. Pellentesque ac metus ac dolor + mattis euismod in id eros. Phasellus sed ornare ligula, sit amet ullamcorper ante. Aenean + fermentum lorem sit amet tortor ultricies, id pulvinar nibh pulvinar. + """ + + def calculate_price(self, input_tokens: int, output_tokens: int) -> Decimal: + price = ( + input_tokens * self.TOKENS_COST['input'] / 1_000_000 + + output_tokens * self.TOKENS_COST['output'] / 1_000_000 + ) + return price.quantize(Decimal('0.01'), rounding='ROUND_UP') + + def save_results(self, content: str, t: timedelta, save: bool = True) -> list[Message]: + msg = Message( + content=content, + content_object=self.store, + elapsed_time=t, + ) + if save: + msg.save() + return [msg] + + def make(self, input_message: Message, save: bool = True) -> list[Message]: + prepare = self._prepare(input_message) + result = ''.join(self._stream(prepare)) + return self._finalize( + input_message, prepare, result, save, output_token_count=len(prepare.output_tokens) + ) + + def make_stream(self, input_message: Message, save: bool = True) -> Iterator[str]: + prepare = self._prepare(input_message) + result = '' + output_token_count = 0 + try: + for token in self._stream(prepare): + result += token + output_token_count += 1 + yield token + finally: + if result: + self._finalize(input_message, prepare, result, save, output_token_count=output_token_count) + return result + + def _prepare(self, input_message: Message) -> PrepareTestData: + start_time = time.time() + info = input_message.info.copy() + input_tokens = self._get_cached_tokens(input_message.content) + output_tokens = self._get_cached_tokens(info.get('cm') or self.BASE_OUTPUT_MESSAGE) + ttft = info.get('ttft', 0.5) + tbt = info.get('tbt', 0.35) + return PrepareTestData(input_tokens, output_tokens, ttft, tbt, start_time) + + def _stream(self, prepare: PrepareTestData) -> Iterator[str]: + time.sleep(prepare.ttft) + for i, token in enumerate(prepare.output_tokens, start=1): + yield token + if i < len(prepare.output_tokens): + time.sleep(prepare.tbt) + + def _finalize( + self, + input_message: Message, + prepare: PrepareTestData, + result: str, + save: bool, + *, + output_token_count: int, + ) -> list[Message]: + process_time = timedelta(seconds=(time.time() - prepare.start_time)) + self.handle_invoice( + input_message.content_object.model, len(prepare.input_tokens), output_token_count + ) + return self.save_results(result, process_time, save) + + @classmethod + def _tokenize(cls, text: str) -> TestTextTokens: + encoding = tiktoken.get_encoding(cls.ENCODING) + return [encoding.decode([token]) for token in encoding.encode(text)] + + @classmethod + def _get_token_cache_key(cls, text: str) -> str: + return hashlib.sha256(f'text_test_model:tokens:{cls.ENCODING}:{text}'.encode('utf-8')).hexdigest() + + @classmethod + def _get_cached_tokens(cls, text: str) -> TestTextTokens: + key = cls._get_token_cache_key(text) + cached = cache.get(key) + if cached is not None: + return cached + tokens = cls._tokenize(text) + cache.set(key, tokens) + return tokens @@ -0,0 +1,89 @@ +import json +from decimal import Decimal +from datetime import timedelta +from unittest.mock import patch + +from core.tests import BaseAuthorizedAPITest +from ml_model.models import ModelCategory, NeuronModel +from tools.chats.models import Chat + + +class TextTestModelAPITest(BaseAuthorizedAPITest): + INPUT_TEXT = ( + 'Lorem ipsum dolor sit amet, consectetur adipiscing elit. Pellentesque ac metus ac dolor mattis euismod in ' + 'id eros. Phasellus sed ornare ligula, sit amet ullamcorper ante. Aenean fermentum lorem sit amet tortor ' + 'ultricies, id pulvinar nibh pulvinar.' + ) + OUTPUT_TEXT = ( + 'Aliquam molestie orci nisl, eget rhoncus nisi varius non. Integer eleifend neque nisi, quis feugiat augue ' + 'malesuada eu. Mauris tincidunt augue id justo ultrices convallis. Nullam a lorem mauris. Duis faucibus est ' + 'mauris, id vestibulum tellus tempor rhoncus.' + ) + TTFT = 0.5 + TBT = 0.35 + + @classmethod + def setup_test_data(cls) -> None: + category = ModelCategory.objects.create(title='Chat-bots', slug='chat-bots') + cls.model = NeuronModel.objects.create( + title='Text Test Model', + slug='text_test_model', + category=category, + ) + cls.chat = Chat.objects.create(title='Text test chat', user=cls.user, model=cls.model) + + @property + def ENDPOINT(self) -> str: + return f'/api/v1/chats/{self.chat.uid}/messages/' + + def _request_data(self) -> dict: + return { + 'content': self.INPUT_TEXT, + 'info': json.dumps( + { + 'ttft': self.TTFT, + 'tbt': self.TBT, + 'cm': self.OUTPUT_TEXT, + } + ), + } + + @staticmethod + def _duration_from_api(value: str) -> Decimal: + hours, minutes, seconds = value.split(':') + total = timedelta( + hours=int(hours), + minutes=int(minutes), + seconds=float(seconds), + ).total_seconds() + return Decimal(str(total)) + + def test_unauthorized_status_code(self) -> None: + response = self.client.post(self.ENDPOINT, data=self._request_data()) + self.assertEqual(response.status_code, 401) + self.assertIn('detail', response.json()) + + @patch('ml_model.services.text_test_model.time.sleep') + def test_generation_time_is_not_more_than_21_seconds(self, _sleep_mock) -> None: + response = self.post(data=self._request_data()) + self.assertEqual(response.status_code, 201, response.json()) + + payload = response.json() + self.assertEqual(len(payload), 2) + self.assertEqual(payload[1]['content'], self.OUTPUT_TEXT) + + elapsed = self._duration_from_api(payload[1]['elapsed_time']) + self.assertLessEqual(elapsed, Decimal('21')) + + @patch('ml_model.services.text_test_model.time.sleep') + def test_billing_uses_rounded_price(self, _sleep_mock) -> None: + expected_charge = Decimal('0.20') + + balance_before = self.user.payment_plan.current_token_balance + response = self.post(data=self._request_data()) + self.assertEqual(response.status_code, 201, response.json()) + + self.user.payment_plan.refresh_from_db() + charged_amount = balance_before - self.user.payment_plan.current_token_balance + + self.assertEqual(charged_amount, expected_charge) @@ -1,10 +1,21 @@ +import importlib from typing import List +from uuid import UUID +from django.http import StreamingHttpResponse +from django.utils.translation import gettext as _ from ninja import Router +from ninja.errors import HttpError from authentication.security import SyncAuthBearer +from messages.models import Message from ml_model.models import NeuronModel from ml_model.schemas import NeuronModelLink +from tools.chats.models import Chat +from tools.chats.schemas import MessageInSchema +from tools.chats.services.sse_chat_stream import SSEChatStreamService +from tools.chats.services.sse_store import SSEStoreService +from tools.chats.tasks import event_stream_task router = Router(auth=SyncAuthBearer(), tags=['chats']) @@ -18,3 +29,48 @@ def get_links(request): ) .order_by('order') ) + + +@router.get('{chat_uid}/messages/{message_uid}/stream', tags=['chats']) +def stream_message_reconnect(request, chat_uid: UUID, message_uid: UUID, offset: int): + store = SSEStoreService(message_uuid=message_uid, user_uuid=request.auth.uid, chat_uuid=chat_uid) + if not store.exists(): + raise HttpError(404, _('Stream not found')) + sse_chat_stream = SSEChatStreamService(store) + return StreamingHttpResponse( + sse_chat_stream.event_stream(request=request, offset=offset), + content_type='text/event-stream', + headers={'Cache-Control': 'no-cache', 'Connection': 'keep-alive', 'X-Accel-Buffering': 'no'}, + ) + + +@router.post('{chat_uid}/messages/stream', tags=['chats']) +def stream_message(request, chat_uid: UUID, body: MessageInSchema): + try: + chat = Chat.objects.select_related('model').get(pk=chat_uid, user=request.auth) + except Chat.DoesNotExist: + raise HttpError(404, _('Chat not found')) + + slug = chat.model.slug + service = getattr(importlib.import_module(f'ml_model.services.{slug}'), f'{slug.title()}') + if not hasattr(service, 'make_stream'): + raise HttpError(501, _('Stream not supported for this model')) + + data = body.model_dump(include={'content', 'file', 'info'}, exclude_unset=True) + input_message = Message.objects.create(content_object=chat, from_model=False, **data) + + store = SSEStoreService(message_uuid=input_message.pk, user_uuid=request.auth.uid, chat_uuid=chat.pk) + if store.exists(): + raise HttpError(409, _('Stream already in progress')) + store.start() + + event_stream_task.delay( + chat_uuid=str(chat.pk), message_uuid=str(input_message.pk), user_uuid=str(request.auth.uid) + ) + + sse_chat_stream = SSEChatStreamService(store) + return StreamingHttpResponse( + sse_chat_stream.event_stream(request=request), + content_type='text/event-stream', + headers={'Cache-Control': 'no-cache', 'Connection': 'keep-alive', 'X-Accel-Buffering': 'no'}, + ) @@ -0,0 +1,82 @@ +import time +from typing import Iterator + +from django.conf import settings +from django.db.models import Exists, OuterRef, Subquery +from django.http import HttpRequest +from django.utils.translation import gettext as _ + +from messages.models import Message +from tools.chats.services.sse_chunk_service import SSEChunkService +from tools.chats.services.sse_store import SSEStoreService + + +class SSEChatStreamService: + HEARTBEAT = ': heartbeat\n\n' + + def __init__(self, store: SSEStoreService) -> None: + self.store = store + + def event_stream(self, request: HttpRequest | None = None, offset: int = 0) -> Iterator[str]: + last_event_id = offset + try: + while self.store.exists(): + if getattr(request, 'closed', False): + return + + chunks = self.store.get_chunks(last_event_id) + if not chunks: + time.sleep(settings.SSE_POLL_INTERVAL) + yield self.HEARTBEAT + continue + + for chunk in chunks: + last_event_id = chunk.event_id + yield chunk.encode() + if chunk.event in ('done', 'error'): + self.store.cleanup() + return + except (GeneratorExit, BrokenPipeError, ConnectionResetError): + return + except Exception as exc: + yield SSEChunkService.error(last_event_id + 1, str(exc)).encode() + self.store.cleanup() + return + + assistant = self._get_assistant_message() + if assistant and assistant.content: + last_event_id += 1 + yield SSEChunkService.done(last_event_id, assistant.content).encode() + self.store.cleanup() + return + + yield SSEChunkService.error(last_event_id + 1, _('Stream timeout')).encode() + self.store.cleanup() + + def _get_assistant_message(self) -> Message | None: + input_qs = Message.objects.filter(pk=self.store.message_uuid, is_deleted=False) + input_created_at = Subquery(input_qs.values('created_at')[:1]) + + return ( + Message.objects.filter( + from_model=True, + is_deleted=False, + content_type=Subquery(input_qs.values('content_type')[:1]), + object_id=Subquery(input_qs.values('object_id')[:1]), + created_at__gt=input_created_at, + ) + .annotate( + has_gap=Exists( + Message.objects.filter( + content_type=OuterRef('content_type'), + object_id=OuterRef('object_id'), + is_deleted=False, + created_at__gt=input_created_at, + created_at__lt=OuterRef('created_at'), + ) + ) + ) + .filter(has_gap=False) + .order_by('created_at') + .first() + ) @@ -0,0 +1,24 @@ +from tools.chats.domain import SSEChunk +from tools.chats.typing import SSEData, SSEEvent + + +class SSEChunkService: + @classmethod + def _chunk(cls, event_id: int, event: SSEEvent, data: SSEData | None = None) -> SSEChunk: + return SSEChunk(event_id=event_id, event=event, data=data or {}) + + @classmethod + def start(cls, event_id: int, message_uuid: str) -> SSEChunk: + return cls._chunk(event_id, 'start', {'message_uuid': message_uuid}) + + @classmethod + def token(cls, event_id: int, content: str) -> SSEChunk: + return cls._chunk(event_id, 'token', {'content': content}) + + @classmethod + def error(cls, event_id: int, content: str) -> SSEChunk: + return cls._chunk(event_id, 'error', {'detail': content}) + + @classmethod + def done(cls, event_id: int, content: str = '') -> SSEChunk: + return cls._chunk(event_id, 'done', {'content': content}) @@ -0,0 +1,59 @@ +import orjson + +from uuid import UUID + +from django.conf import settings +from django.core.cache import caches + +from tools.chats.domain import SSEChunk + + +class SSEStoreService: + def __init__(self, message_uuid: UUID, user_uuid: UUID, chat_uuid: UUID) -> None: + self.message_uuid = message_uuid + self.user_uuid = user_uuid + self.chat_uuid = chat_uuid + self.redis_client = self._get_redis_client() + + @staticmethod + def _get_redis_client(): + return caches['default']._cache.get_client(write=True) + + def _get_cache_key(self) -> str: + return f'sse:tokens:{self.user_uuid}:{self.chat_uuid}:{self.message_uuid}' + + def start(self) -> None: + pipe = self.redis_client.pipeline() + pipe.rpush( + self._get_cache_key(), + orjson.dumps( + { + 'event_id': 0, + 'event': 'pending', + 'data': {}, + } + ), + ) + pipe.expire(self._get_cache_key(), settings.SSE_STREAM_TTL) + pipe.execute() + + def push(self, sse_chunk: SSEChunk, ttl: int | None = None) -> None: + pipe = self.redis_client.pipeline() + pipe.rpush(self._get_cache_key(), orjson.dumps(sse_chunk.to_dict())) + pipe.expire(self._get_cache_key(), ttl if ttl is not None else settings.SSE_STREAM_TTL) + pipe.execute() + + def exists(self) -> bool: + return bool(self.redis_client.exists(self._get_cache_key())) + + def get_chunks(self, offset: int = 0) -> list[SSEChunk]: + return [ + SSEChunk(**orjson.loads(raw)) + for raw in self.redis_client.lrange(self._get_cache_key(), offset + 1, -1) + ] + + def delete_stream(self) -> None: + self.redis_client.delete(self._get_cache_key()) + + def cleanup(self) -> None: + self.delete_stream() @@ -0,0 +1,19 @@ +import orjson + +from dataclasses import asdict, dataclass + +from tools.chats.typing import SSEData, SSEEvent + + +@dataclass +class SSEChunk: + event_id: int + event: SSEEvent + data: SSEData + + def encode(self): + payload = orjson.dumps(self.data, default=str).decode() + return f'id: {self.event_id}\nevent: {self.event}\ndata: {payload}\n\n' + + def to_dict(self): + return asdict(self) @@ -0,0 +1,75 @@ +import orjson + +from django.utils.translation import gettext as _ +from ninja import ModelSchema, UploadedFile +from pydantic import field_serializer, field_validator, model_validator + +from messages.models import Message + + +class MessageInSchema(ModelSchema): + file: UploadedFile | None = None + + class Meta: + model = Message + fields = ['content', 'file', 'info'] + fields_optional = '__all__' + + @model_validator(mode='after') + def validate_file_size(self): + max_mb_size = 50 + if self.file and hasattr(self.file, 'size') and self.file.size > (max_mb_size << 10 << 10): + raise ValueError( + _('The file size cannot exceed %(max_mb_size)d MB') % {'max_mb_size': max_mb_size} + ) + return self + + @field_validator('info', mode='before', check_fields=False) + @classmethod + def validate_info(cls, obj): + if isinstance(obj, str): + try: + return orjson.loads(obj) + except orjson.JSONDecodeError: + raise ValueError(_('Invalid info payload')) + return obj + + +class MessageSchema(ModelSchema): + model: str | None = None + + class Meta: + model = Message + fields = [ + 'uid', + 'content', + 'file', + 'from_model', + 'created_at', + 'elapsed_time', + 'is_favourite', + 'is_sent', + 'info', + ] + + @staticmethod + def resolve_model(obj: Message) -> str | None: + try: + return obj.content_object.model.slug + except AttributeError: + return None + + @field_serializer('content', check_fields=False) + def serialize_content(self, value: str | None) -> str | None: + if value: + return value.replace('\\n', '\n') + return value + + @field_serializer('file', check_fields=False) + def serialize_file(self, value) -> str | None: + if value is None: + return None + if isinstance(value, str): + return value + return value.url + @@ -0,0 +1,47 @@ +import importlib + +from celery import shared_task +from django.conf import settings +from messages.models import Message +from tools.chats.models import Chat +from tools.chats.services.sse_chunk_service import SSEChunkService +from tools.chats.services.sse_store import SSEStoreService + + +@shared_task(soft_time_limit=570, time_limit=600) +def event_stream_task(chat_uuid: str, message_uuid: str, user_uuid: str) -> None: + store = SSEStoreService(message_uuid, user_uuid, chat_uuid) + stream = None + event_id = 0 + + try: + chat = Chat.objects.select_related('model').get(pk=chat_uuid) + message = Message.objects.get(pk=message_uuid) + + slug = chat.model.slug + service_cls = getattr(importlib.import_module(f'ml_model.services.{slug}'), f'{slug.title()}') + + event_id += 1 + store.push(SSEChunkService.start(event_id, message_uuid)) + + stream = service_cls(chat).make_stream(message) + while True: + token = next(stream) + if not token: + continue + event_id += 1 + store.push(SSEChunkService.token(event_id, token)) + except StopIteration as exc: + event_id += 1 + store.push(SSEChunkService.done(event_id, exc.value or ''), ttl=settings.SSE_DONE_STREAM_TTL) + except (Chat.DoesNotExist, Message.DoesNotExist) as exc: + store.push(SSEChunkService.error(event_id + 1, str(exc)), ttl=settings.SSE_DONE_STREAM_TTL) + except Exception as exc: + if event_id < 2: + message.is_sent = False + message.save(update_fields=['is_sent']) + event_id += 1 + store.push(SSEChunkService.error(event_id, str(exc)), ttl=settings.SSE_DONE_STREAM_TTL) + finally: + if stream is not None: + stream.close() @@ -0,0 +1,4 @@ +from typing import Any, Literal + +type SSEEvent = Literal['pending', 'start', 'token', 'error', 'done'] +type SSEData = dict[str, Any] @@ -79,6 +79,9 @@ MAX_THREADS=3 # REDIS REDIS_HOST=cache-mdb REDIS_PORT=6379 +SSE_STREAM_TTL=600 +SSE_DONE_STREAM_TTL=30 +SSE_POLL_INTERVAL=0.2 # DEBUG USER DJANGO_SUPERUSER_EMAIL=example@root.ru