@@ -257,10 +257,15 @@ 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', 15) +SSE_POLL_INTERVAL = env.float('SSE_POLL_INTERVAL', 0.2) # Celery CELERY_BROKER_URL = env.str('CELERY_BROKER_URL', 'redis://celery-mdb:6379/0') CELERY_RESULT_BACKEND = env.str('CELERY_RESULT_BACKEND', 'redis://celery-mdb:6379/0') +CELERY_WORKER_SOFT_SHUTDOWN_TIMEOUT = env.int('CELERY_WORKER_SOFT_SHUTDOWN_TIMEOUT', 630) +CELERY_WORKER_PREFETCH_MULTIPLIER = env.int('CELERY_WORKER_PREFETCH_MULTIPLIER', 1) CELERY_ACCEPT_CONTENT = ['json', 'application/x-python-serialize', 'pickle'] CELERY_RESULT_SERIALIZER = 'pickle' @@ -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-ключа" @@ -1,5 +1,3 @@ -import sys - from drf_spectacular.utils import extend_schema from rest_framework.pagination import LimitOffsetPagination from rest_framework.permissions import IsAuthenticated @@ -9,7 +7,6 @@ from rest_framework.views import APIView from messages.models import Message from messages.serializers import MessageSerializer from ml_model.models import ModelParameter -from ml_model.services.base import SimpleService from tools.chats.models import Chat @@ -40,9 +37,7 @@ class MessagesAPIView(APIView): chat = Chat.objects.get(pk=chat_uid) info = {} if chat.model: - service: type[SimpleService] = getattr( - sys.modules['ml_model.services'], f'{chat.model.title}' - ) + service = chat.model.service missing_info = {} for p in chat.model.parameters.difference( ModelParameter.objects.filter(model=chat.model, key__in=info.keys()) @@ -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: ... @@ -1,17 +1,31 @@ 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 SimpleService +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(SimpleService): + +class Text_Test_Model(StreamSimpleService): TOKENS_COST = { 'input': Decimal('1000'), # per 1 million input-tokens 'output': Decimal('2500'), # per 1 million output-tokens @@ -43,24 +57,59 @@ class Text_Test_Model(SimpleService): return [msg] def make(self, input_message: Message, save: bool = True) -> list[Message]: - info = input_message.info.copy() + 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) - time.sleep(ttft) - result = '' - for i, token in enumerate(output_tokens, start=1): - result += token - if i < len(output_tokens): - time.sleep(tbt) - process_time = timedelta(seconds=(time.time() - start_time)) - self.handle_invoice(input_message.content_object.model, len(input_tokens), len(output_tokens)) + 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) -> list[str]: + def _tokenize(cls, text: str) -> TestTextTokens: encoding = tiktoken.get_encoding(cls.ENCODING) return [encoding.decode([token]) for token in encoding.encode(text)] @@ -69,7 +118,7 @@ class Text_Test_Model(SimpleService): return hashlib.sha256(f'text_test_model:tokens:{cls.ENCODING}:{text}'.encode('utf-8')).hexdigest() @classmethod - def _get_cached_tokens(cls, text: str) -> list[str]: + def _get_cached_tokens(cls, text: str) -> TestTextTokens: key = cls._get_token_cache_key(text) cached = cache.get(key) if cached is not None: @@ -77,5 +126,3 @@ class Text_Test_Model(SimpleService): tokens = cls._tokenize(text) cache.set(key, tokens) return tokens - - @@ -1,3 +1,4 @@ +import importlib from uuid import uuid4 from django.contrib.contenttypes.fields import GenericForeignKey @@ -140,6 +141,14 @@ class NeuronModel(BaseModel, OrderedModel): def blocked(self) -> bool: return not self.active + @property + def service(self) -> 'SimpleService': # noqa: F821 + return getattr(importlib.import_module(f'ml_model.services.{self.slug}'), self.slug.replace('-', '').title()) + + @property + def streaming(self) -> bool: + return hasattr(self.service, 'make_stream') + def __str__(self): return self.title @@ -59,6 +59,7 @@ class NeuronModelSerializer(serializers.ModelSerializer): tags = ModelTagSerializer(many=True) instruction = ModelInstructionSerializer() blocked = serializers.BooleanField() + streaming = serializers.BooleanField(read_only=True) class Meta: model = NeuronModel @@ -83,10 +84,7 @@ class PublicNeuronModelSerializer(NeuronModelSerializer): class Meta: model = NeuronModel - fields = ( - 'title', 'description', 'slug', 'blocked', 'versions', - 'inputs', 'parameters' - ) + fields = ('title', 'description', 'slug', 'blocked', 'versions', 'inputs', 'parameters') class NeuronModelsSerializer(serializers.ModelSerializer): @@ -1,10 +1,20 @@ 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 +28,46 @@ def get_links(request): ) .order_by('order') ) + + +@router.get('{chat_uid}/messages/stream', tags=['chats']) +def stream_message_reconnect(request, chat_uid: UUID, offset: int = 0): + store = SSEStoreService(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')) + + if not chat.model.streaming: + raise HttpError(501, _('Stream not supported for this model')) + + store = SSEStoreService(user_uuid=request.auth.uid, chat_uuid=chat_uid) + if store.exists(): + raise HttpError(409, _('Stream already in progress')) + + data = body.model_dump(include={'content', 'file', 'info'}, exclude_unset=True) + input_message = Message.objects.create(content_object=chat, from_model=False, **data) + 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,85 @@ +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 + message_uuid = self.store.get_message_uuid() + 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: + if chunk.event == 'start': + message_uuid = chunk.data.get('message_uuid') or message_uuid + 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: + return + + assistant = self._get_assistant_message(message_uuid) + 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() + + def _get_assistant_message(self, message_uuid) -> Message | None: + if not message_uuid: + return None + + input_qs = Message.objects.filter(pk=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,68 @@ +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, user_uuid: UUID, chat_uuid: UUID) -> None: + 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}' + + 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_message_uuid(self) -> UUID | None: + raw = self.redis_client.lindex(self._get_cache_key(), 1) + if not raw: + return None + chunk = orjson.loads(raw) + if chunk.get('event') != 'start': + return None + uid = chunk.get('data', {}).get('message_uuid') + return UUID(str(uid)) if uid else None + + 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,151 @@ +import os +import random +import resource +import sys +import time +from concurrent.futures import ThreadPoolExecutor, as_completed +from dataclasses import dataclass +from unittest.mock import patch + +import orjson +from cacheops import invalidate_all +from django.db import close_old_connections, connections +from django.test import Client, TransactionTestCase +from rest_framework_simplejwt.tokens import RefreshToken + +from authentication.models import CustomUserModel +from ml_model.models import ModelCategory, NeuronModel +from ml_model.services.text_test_model import Text_Test_Model +from tools.chats.models import Chat +from tools.chats.tasks import event_stream_task + +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.' +) + +MAX_POOL_WORKERS = 10 + + +@dataclass(frozen=True) +class LoadScenario: + clients: int + delay_min: float + delay_max: float + ttft: float + tbt: float + + +LOAD_SCENARIOS = { + 5: LoadScenario(5, 0.5, 5.0, 0.45, 0.15), + 50: LoadScenario(50, 0.7, 7.0, 0.55, 0.35), + 100: LoadScenario(100, 0.25, 10.0, 0.5, 0.35), +} + + +def _out(message: str) -> None: + sys.__stdout__.write(f'{message}\n') + sys.__stdout__.flush() + + +def _cpu_time_s() -> float: + usage = resource.getrusage(resource.RUSAGE_SELF) + return usage.ru_utime + usage.ru_stime + + +def _ram_mb() -> float: + try: + with open('/proc/self/status', encoding='utf-8') as status: + for line in status: + if line.startswith('VmRSS:'): + return int(line.split()[1]) / 1024 + except OSError: + pass + rss = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss + return rss / 1024 if sys.platform != 'darwin' else rss / 1024 / 1024 + + +class SSEStreamLoadTest(TransactionTestCase): + scenario: LoadScenario + + def setUp(self) -> None: + invalidate_all() + self.user = CustomUserModel.objects.create_user(email='sse-load@test.test', password='test') + self.access_token = str(RefreshToken.for_user(self.user).access_token) + category = ModelCategory.objects.create(title='Chat-bots', slug='chat-bots') + model = NeuronModel.objects.create( + title='Text Test Model', slug='text_test_model', category=category + ) + self.chat = Chat.objects.create(title='SSE load chat', user=self.user, model=model) + self.stream_url = f'/api/v1/api/chats/{self.chat.uid}/messages/stream' + Text_Test_Model._get_cached_tokens(INPUT_TEXT) + Text_Test_Model._get_cached_tokens(OUTPUT_TEXT) + + def _run_client(self, _client_id: int) -> None: + close_old_connections() + try: + time.sleep(random.uniform(self.scenario.delay_min, self.scenario.delay_max)) + response = Client().post( + self.stream_url, + data=orjson.dumps( + { + 'content': INPUT_TEXT, + 'info': {'ttft': self.scenario.ttft, 'tbt': self.scenario.tbt, 'cm': OUTPUT_TEXT}, + } + ), + content_type='application/json', + HTTP_AUTHORIZATION=f'Bearer {self.access_token}', + ) + if response.status_code != 200: + raise RuntimeError(f'HTTP {response.status_code}') + for _ in response.streaming_content: + pass + finally: + connections.close_all() + + def _run_load_test(self, scenario: LoadScenario) -> None: + self.scenario = scenario + cpu_cores = os.cpu_count() or 1 + cpu_start = _cpu_time_s() + ram_start = _ram_mb() + wall_start = time.perf_counter() + errors = 0 + pool_workers = min(scenario.clients, MAX_POOL_WORKERS) + + with patch('tools.chats.routes.v1.event_stream_task.delay') as delay_mock: + delay_mock.side_effect = lambda **kwargs: event_stream_task(**kwargs) + with ThreadPoolExecutor(max_workers=pool_workers) as pool: + futures = [pool.submit(self._run_client, i) for i in range(scenario.clients)] + for future in as_completed(futures): + try: + future.result() + except Exception: + errors += 1 + + connections.close_all() + wall = time.perf_counter() - wall_start + cpu = _cpu_time_s() - cpu_start + ram = _ram_mb() + cpu_percent = cpu / wall / cpu_cores * 100 if wall else 0 + + self.assertEqual(errors, 0) + + _out( + f'cpu={cpu:.2f}s ({cpu_percent:.1f}% of {cpu_cores} cores), ' + f'ram={ram:.1f}MB (delta {ram - ram_start:+.1f}MB)' + ) + + def test_sse_load_5_clients(self) -> None: + self._run_load_test(LOAD_SCENARIOS[5]) + + def test_sse_load_50_clients(self) -> None: + self._run_load_test(LOAD_SCENARIOS[50]) + + def test_sse_load_100_clients(self) -> None: + self._run_load_test(LOAD_SCENARIOS[100]) @@ -0,0 +1,93 @@ +import sys +import time +from unittest.mock import patch + +import orjson +from cacheops import invalidate_all +from django.test import Client, TestCase +from rest_framework_simplejwt.tokens import RefreshToken + +from authentication.models import CustomUserModel +from ml_model.models import ModelCategory, NeuronModel +from ml_model.services.text_test_model import Text_Test_Model +from tools.chats.models import Chat +from tools.chats.tasks import event_stream_task + + +def test_out(message: str) -> None: + sys.__stdout__.write(f'{message}\n') + sys.__stdout__.flush() + + +class SSEStreamAPITest(TestCase): + TTFT = 0.5 + TBT = 0.35 + MIN_CHUNK_TRANSFER = 0.001 + MAX_CHUNK_TRANSFER = 0.08 + + 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.' + ) + + def setUp(self) -> None: + invalidate_all() + self.client = Client() + self.user = CustomUserModel.objects.create_user(email='sse-test@test.test', password='test') + self.access_token = str(RefreshToken.for_user(self.user).access_token) + category = ModelCategory.objects.create(title='Chat-bots', slug='chat-bots') + model = NeuronModel.objects.create( + title='Text Test Model', + slug='text_test_model', + category=category, + ) + self.chat = Chat.objects.create(title='SSE test chat', user=self.user, model=model) + Text_Test_Model._get_cached_tokens(self.INPUT_TEXT) + Text_Test_Model._get_cached_tokens(self.OUTPUT_TEXT) + + @patch('tools.chats.routes.v1.event_stream_task.delay') + def test_sse_stream_duration(self, delay_mock) -> None: + delay_mock.side_effect = lambda **kwargs: event_stream_task(**kwargs) + + started = time.perf_counter() + response = self.client.post( + f'/api/v1/api/chats/{self.chat.uid}/messages/stream', + data=orjson.dumps( + { + 'content': self.INPUT_TEXT, + 'info': {'ttft': self.TTFT, 'tbt': self.TBT, 'cm': self.OUTPUT_TEXT}, + } + ), + content_type='application/json', + HTTP_AUTHORIZATION=f'Bearer {self.access_token}', + ) + self.assertEqual(response.status_code, 200) + + token_count = 0 + for chunk in response.streaming_content: + text = chunk.decode() if isinstance(chunk, bytes) else chunk + token_count += text.count('event: token') + + duration = time.perf_counter() - started + test_out(f'SSE total: {duration:.3f}s') + + model_time = self.TTFT + (token_count - 1) * self.TBT + min_duration = model_time + token_count * self.MIN_CHUNK_TRANSFER + max_duration = model_time + token_count * self.MAX_CHUNK_TRANSFER + + self.assertGreaterEqual( + duration, + min_duration, + f'{duration:.3f}s is below min {min_duration:.3f}s ({token_count} tokens)', + ) + self.assertLessEqual( + duration, + max_duration, + f'{duration:.3f}s is above max {max_duration:.3f}s ({token_count} tokens)', + ) @@ -1,10 +1,11 @@ import json -from decimal import Decimal from datetime import timedelta +from decimal import Decimal from unittest.mock import patch from core.tests import BaseAuthorizedAPITest from ml_model.models import ModelCategory, NeuronModel +from ml_model.services.text_test_model import Text_Test_Model from tools.chats.models import Chat @@ -32,6 +33,11 @@ class TextTestModelAPITest(BaseAuthorizedAPITest): ) cls.chat = Chat.objects.create(title='Text test chat', user=cls.user, model=cls.model) + def setUp(self) -> None: + super().setUp() + Text_Test_Model._get_cached_tokens(self.INPUT_TEXT) + Text_Test_Model._get_cached_tokens(self.OUTPUT_TEXT) + @property def ENDPOINT(self) -> str: return f'/api/v1/chats/{self.chat.uid}/messages/' @@ -63,8 +69,7 @@ class TextTestModelAPITest(BaseAuthorizedAPITest): 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: + def test_generation_time_is_not_more_than_21_seconds(self) -> None: response = self.post(data=self._request_data()) self.assertEqual(response.status_code, 201, response.json()) @@ -73,7 +78,7 @@ class TextTestModelAPITest(BaseAuthorizedAPITest): self.assertEqual(payload[1]['content'], self.OUTPUT_TEXT) elapsed = self._duration_from_api(payload[1]['elapsed_time']) - self.assertLessEqual(elapsed, Decimal('21')) + self.assertLess(elapsed, Decimal('21')) @patch('ml_model.services.text_test_model.time.sleep') def test_billing_uses_rounded_price(self, _sleep_mock) -> None: @@ -1,5 +1,4 @@ import logging -import sys from django.utils.translation import gettext_lazy as _ from drf_spectacular.utils import OpenApiParameter, extend_schema @@ -36,7 +35,6 @@ from ml_model.exceptions import ( TemplateUnknownException, UnrecognizedFileError, ) -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 @@ -156,9 +154,7 @@ class MessagesAPIView(APIView): if serializer.is_valid(): chat = Chat.objects.get(pk=chat_uid) info = serializer.validated_data.pop('info', {}) - service: type[SimpleService] = getattr( - sys.modules['ml_model.services'], f'{chat.model.slug.title()}' - ) + service = chat.model.service input_message = Message.objects.create( **serializer.validated_data, info=info, @@ -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,42 @@ +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(user_uuid=user_uuid, chat_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) + + event_id += 1 + store.push(SSEChunkService.start(event_id, message_uuid)) + + stream = chat.model.service(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] @@ -1,5 +1,4 @@ import logging -import sys from django.utils.translation import gettext_lazy as _ from drf_spectacular.utils import OpenApiParameter, extend_schema @@ -31,7 +30,6 @@ from ml_model.exceptions import ( UnsupportedSize, ) from ml_model.models import NeuronModel -from ml_model.services.base import SimpleService from payments.exceptions.insufficient_balance import InsufficientBalance from .models import Audio, Image, Video, VoiceClone, Voice, Preset @@ -154,10 +152,7 @@ class MediaAPIView(APIView): }, ) info = serializer.validated_data.pop('info', {}) - service: type[SimpleService] = getattr( - sys.modules['ml_model.services'], - f'{gallery.model.slug.replace("-", "").title()}', - ) + service = gallery.model.service input_message = Message.objects.create( **serializer.validated_data, info=info, @@ -1,5 +1,4 @@ import logging -import sys from django.core.exceptions import ValidationError from django.utils.translation import gettext_lazy as _ @@ -19,7 +18,6 @@ from ml_model.exceptions import InvalidParameterError from ml_model.models import NeuronModel from ml_model.selectors.ml_models_selector import NeuronModelSelector from ml_model.serializers import PublicNeuronModelSerializer -from ml_model.services.base import SimpleService from payments.exceptions.insufficient_balance import InsufficientBalance from tools.public_api.models import APIKey, APIStore from tools.public_api.permissions import HasAPIKey @@ -77,7 +75,7 @@ class BaseGenerationView(APIView): {'detail': _('The request must not be empty')}, status=HTTP_400_BAD_REQUEST, ) - service: type[SimpleService] = getattr(sys.modules['ml_model.services'], f'{model.slug.title()}') + service = model.service info = serializer.validated_data.pop('info', {}) input_message = Message.objects.create( **serializer.validated_data, @@ -52,6 +52,8 @@ MINIO_SECRET_KEY=testtest # CELERY CELERY_BROKER_URL=redis://cache-mdb:6379/0 CELERY_RESULT_BACKEND=redis://cache-mdb:6379/0 +CELERY_WORKER_SOFT_SHUTDOWN_TIMEOUT=630 +CELERY_WORKER_PREFETCH_MULTIPLIER=1 # EMAIL # For free hosts - https://www.wpoven.com/tools/free-smtp-server-for-testing @@ -79,6 +81,9 @@ MAX_THREADS=3 # REDIS REDIS_HOST=cache-mdb REDIS_PORT=6379 +SSE_STREAM_TTL=600 +SSE_DONE_STREAM_TTL=15 +SSE_POLL_INTERVAL=0.2 # DEBUG USER DJANGO_SUPERUSER_EMAIL=example@root.ru @@ -110,7 +110,8 @@ services: - C_FORCE_ROOT=true - RELEASE - ENVIRONMENT - command: celery -A backend worker -l INFO --concurrency 3 + command: ["celery", "-A", "backend", "worker", "-l", "INFO", "--concurrency", "3"] + stop_grace_period: 10m deploy: replicas: 1 <<: [ *default-deploy ]