@@ -258,7 +258,7 @@ SPECTACULAR_SETTINGS = { 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_DONE_STREAM_TTL = env.int('SSE_DONE_STREAM_TTL', 15) SSE_POLL_INTERVAL = env.float('SSE_POLL_INTERVAL', 0.2) # Celery @@ -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,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}'), f'{self.slug.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,4 +1,3 @@ -import importlib from typing import List from uuid import UUID @@ -51,9 +50,7 @@ def stream_message(request, chat_uid: UUID, body: MessageInSchema): 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'): + if not chat.model.streaming: raise HttpError(501, _('Stream not supported for this model')) data = body.model_dump(include={'content', 'file', 'info'}, exclude_unset=True) @@ -38,9 +38,7 @@ class SSEChatStreamService: return except (GeneratorExit, BrokenPipeError, ConnectionResetError): return - except Exception as exc: - yield SSEChunkService.error(last_event_id + 1, str(exc)).encode() - self.store.cleanup() + except Exception: return assistant = self._get_assistant_message() @@ -51,7 +49,6 @@ class SSEChatStreamService: 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) @@ -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, @@ -1,5 +1,3 @@ -import importlib - from celery import shared_task from django.conf import settings from messages.models import Message @@ -18,13 +16,10 @@ def event_stream_task(chat_uuid: str, message_uuid: str, user_uuid: str) -> None 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) + stream = chat.model.service(chat).make_stream(message) while True: token = next(stream) if not token: @@ -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, @@ -80,7 +80,7 @@ MAX_THREADS=3 REDIS_HOST=cache-mdb REDIS_PORT=6379 SSE_STREAM_TTL=600 -SSE_DONE_STREAM_TTL=30 +SSE_DONE_STREAM_TTL=15 SSE_POLL_INTERVAL=0.2 # DEBUG USER