@@ -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,14 +15,17 @@ 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) +public_api = NinjaAPI( + title='PUBLIC AIR API', urls_namespace='public-api-1.0.0', parser=MultiContentTypeParser(), docs_url=None +) api.add_router('users/', 'users.routes.v1.router') api.add_router('chats/', 'tools.chats.routes.v1.router') @@ -35,6 +38,8 @@ compatibility_api.add_router('ml_model/', 'ml_model.routes.v1.router') compatibility_api_v2.add_router('auth/', 'authentication.routes.v2.router') +public_api.add_router('', 'tools.public_api.routes.v1.router') + class Status(Schema): status: Literal['ok', 'dead'] @@ -97,6 +102,7 @@ urlpatterns = ( path('api/v1/api/', api.urls), path('api/v1/', compatibility_api.urls), path('api/v1/v2/', compatibility_api_v2.urls), + path('public/', public_api.urls), ] + static(settings.STATIC_URL, document_root=settings.STATIC_ROOT) + public_urlpatterns @@ -122,3 +128,4 @@ if settings.DEBUG: api.docs_url = '/docs' compatibility_api.docs_url = '/docs' compatibility_api_v2.docs_url = '/docs' + public_api.docs_url = '/docs' @@ -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()) @@ -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 @@ -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', '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', '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,79 @@ +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() + + +class PublicSSEStoreService(SSEStoreService): + def __init__(self, message_uuid: UUID, user_uuid: UUID): + self.user_uuid = user_uuid + self.message_uuid = message_uuid + self.redis_client = self._get_redis_client() + + def _get_cache_key(self) -> str: + return f'sse:tokens:{self.user_uuid}:{self.message_uuid}:public' + \ No newline at end of file @@ -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)', + ) @@ -0,0 +1,94 @@ +import json +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 + + +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) + + 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/' + + 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()) + + 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()) + + 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.assertLess(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,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,95 @@ +from decimal import Decimal + +from celery import shared_task +from django.conf import settings +from django.db.models import F, Value +from django.db.models.functions import Greatest + +from messages.models import Message +from ml_model.models import NeuronModel +from ml_model.services.base import StreamSimpleService +from tools.chats.models import Chat +from tools.chats.services.sse_chunk_service import SSEChunkService +from tools.chats.services.sse_store import PublicSSEStoreService, SSEStoreService +from tools.public_api.models import APIKey, APIStore + + +def _run_stream(store: SSEStoreService, message_uuid: str, service: StreamSimpleService, message): + stream = None + event_id = 0 + + try: + event_id += 1 + store.push(SSEChunkService.start(event_id, message_uuid)) + + stream = service.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 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() + + +@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) + + chat = Chat.objects.select_related('model').get(pk=chat_uuid) + message = Message.objects.get(pk=message_uuid) + + service = chat.model.service(chat) + + _run_stream(store, message_uuid, service, message) + + +@shared_task(soft_time_limit=570, time_limit=600) +def public_event_stream_task( + start_user_balance: Decimal, + message_uuid: str, + user_uuid: str, + model_slug: str, + api_key_uuid: str, + debit_api_key_limit: bool, +): + store = PublicSSEStoreService(user_uuid=user_uuid, message_uuid=message_uuid) + + message = Message.objects.get(pk=message_uuid) + api_store = APIStore.objects.select_related( + 'user', + 'user__payment_plan', + 'user__payment_plan__plan', + 'user__business_account', + 'user__business_account__group', + 'user__business_account__parent_company', + 'user__business_account__parent_company__user__payment_plan', + 'user__business_account__parent_company__user__payment_plan__plan', + ).get(pk=message.object_id) + model = NeuronModel.objects.get(slug=model_slug) + + message.content_object = api_store + message.content_object.model = model + + service = model.service(api_store) + + try: + _run_stream(store, message_uuid, service, message) + finally: + if debit_api_key_limit: + spent = start_user_balance - api_store.user.balance + if spent > 0: + APIKey.objects.filter(pk=api_key_uuid).update( + token_limit=Greatest(F('token_limit') - spent, Value(Decimal('0'))), + ) @@ -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, @@ -0,0 +1,172 @@ +import base64 +from io import BytesIO + +import filetype +from django.db.models import Q +from django.core.files.uploadedfile import InMemoryUploadedFile +from django.utils.translation import gettext as _ +from ninja.errors import HttpError + +from ml_model.models import NeuronModel +from tools.chats.schemas import MessageInSchema +from tools.public_api.services.openai_errors import OpenAIErrorService +from tools.public_api.services.openai_stream import OpenAIStreamService + + +def _parse_body(body: dict) -> MessageInSchema: + raw = body.get('input') + if raw is None: + raise HttpError(400, _('Missing required parameter: input')) + file, lines = None, [] + if isinstance(raw, str): + content = raw.strip() + elif not isinstance(raw, list): + raise HttpError(400, _('Invalid input payload')) + else: + for item in raw: + if not isinstance(item, dict): + continue + role = item.get('role', 'user').capitalize() + ic = item.get('content') + if isinstance(ic, str): + if t := ic.strip(): + lines.append(f'[{role}] {t}') + continue + for part in ic or []: + if not isinstance(part, dict): + continue + pt = part.get('type') + if pt in ('input_text', 'text') and (t := part.get('text')) and (t := str(t).strip()): + lines.append(f'[{role}] {t}') + elif pt in ('input_image', 'image_url') and file is None: + url = part.get('image_url') + url = url.get('url', '') if isinstance(url, dict) else str(url or '') + if 'base64' in url: + buf = BytesIO(base64.b64decode(url.split('base64', 1)[-1].lstrip(','))) + kind = filetype.guess(buf.read(20)) + buf.seek(0) + file = InMemoryUploadedFile( + buf, + 'file', + f'api-file.{kind.extension if kind else "bin"}', + kind.mime if kind else 'application/octet-stream', + buf.getbuffer().nbytes, + None, + ) + content = '\n'.join(lines).strip() + if not content and not file: + raise HttpError(400, _('The request must not be empty')) + kwargs = {'content': content or '[User]'} + if file: + kwargs['file'] = file + if isinstance(md := body.get('metadata'), dict) and md: + info = dict(md) + for k in ('ttft', 'tbt'): + if isinstance(info.get(k), str): + try: + info[k] = float(info[k]) + except ValueError: + pass + kwargs['info'] = info + return MessageInSchema(**kwargs) + + +def _resolve_model(model_ref: str) -> NeuronModel: + try: + return ( + NeuronModel.objects.filter( + Q(model_modelversions__slug=model_ref) | Q(slug=model_ref), + category__slug='chat-bots', + ) + .distinct() + .get() + ) + except NeuronModel.DoesNotExist: + raise HttpError(404, _('Model not found')) + + +def _public_user_uuid(request): + from tools.public_api.routes.v1 import _get_api_key + + return _get_api_key(request, check_usage_limit=False).user.pk + + +def _to_openai( + response, + model_ref: str, + *, + request, + message_uuid: str | None = None, + starting_after: int = 0, +): + # Django 6: .streaming_content yields bytes; wrap the raw str iterator instead. + response.streaming_content = OpenAIStreamService.translate( + response._iterator, + model_ref, + message_uuid=message_uuid, + starting_after=starting_after, + user_uuid=_public_user_uuid(request), + ) + return response + + +@OpenAIErrorService.view +def openai_responses_stream(request, body: dict): + from tools.public_api.routes.v1 import public_stream_message + + if not body.get('stream'): + raise HttpError(400, _('Only streaming is supported')) + if not (model_ref := body.get('model')): + raise HttpError(400, _('You must provide a model parameter')) + model = _resolve_model(model_ref) + return _to_openai(public_stream_message(request, model.slug, _parse_body(body)), model_ref, request=request) + + +@OpenAIErrorService.view +def openai_responses_stream_reconnect( + request, + response_id: str | None = None, + *, + stream: bool = True, + starting_after: int = 0, +): + from tools.public_api.routes.v1 import public_stream_message_reconnect + + if not stream: + raise HttpError(400, _('Only streaming reconnect is supported')) + if not (message_uuid := (response_id or request.GET.get('response_id', '')).strip()): + raise HttpError(400, _('message_uuid is not provided')) + response = public_stream_message_reconnect( + request, + message_uuid, + offset=OpenAIStreamService.public_offset(starting_after), + ) + return _to_openai( + response, '', message_uuid=message_uuid, starting_after=starting_after, request=request + ) + + +from tools.public_api.routes import v1 as public_v1_routes + + +def _register_reconnect_get(path: str, *, with_response_id: bool): + if with_response_id: + + @public_v1_routes.router.get(path, tags=['openai/responses']) + def handler(request, response_id: str, stream: bool = True, starting_after: int = 0): + return openai_responses_stream_reconnect( + request, response_id, stream=stream, starting_after=starting_after + ) + else: + + @public_v1_routes.router.get(path, tags=['openai/responses']) + def handler(request, stream: bool = True, starting_after: int = 0): + return openai_responses_stream_reconnect(request, stream=stream, starting_after=starting_after) + + +for with_rid, paths in ( + (False, ('openai/v1/responses', 'openai/responses')), + (True, ('openai/v1/responses/{response_id}', 'openai/responses/{response_id}')), +): + for path in paths: + _register_reconnect_get(path, with_response_id=with_rid) @@ -0,0 +1,141 @@ +from datetime import date + +from django.http import StreamingHttpResponse +from ninja import Body, Router +from ninja.errors import HttpError + +from django.utils.translation import gettext as _ + +from messages.models import Message +from ml_model.selectors.ml_models_selector import NeuronModelSelector + +from tools.chats.schemas import MessageInSchema +from tools.chats.services.sse_chat_stream import SSEChatStreamService +from tools.chats.services.sse_store import PublicSSEStoreService +from tools.chats.tasks import public_event_stream_task + +from tools.public_api.models import APIKey, APIStore + +router = Router(auth=None, tags=['public']) + + +def _get_api_key( + request, + select_related: list[str] = None, + prefetch_related: list[str] = None, + *, + check_usage_limit: bool = True, +): + raw_api_key = request.headers.get('Authorization', '') + if not raw_api_key: + raise HttpError(401, _('No API Key in Authorization header')) + + if (split_api_key := raw_api_key.split())[0] == 'Bearer': + raw_api_key = split_api_key[-1] + + api_key = ( + APIKey.objects.select_related('user', *(select_related or [])) + .prefetch_related(*(prefetch_related or [])) + .filter(key=raw_api_key, is_deleted=False) + ).first() + + if not api_key: + raise HttpError(404, _('API key not found')) + if check_usage_limit: + if api_key.expires_at and api_key.expires_at < date.today(): + raise HttpError(401, _('API key expired')) + if api_key.token_limit is not None and api_key.token_limit < 1: + raise HttpError(403, _('API key limit exceeded')) + + return api_key + + +@router.get('text/{message_uuid}/stream/reconnect', tags=['public/text']) +def public_stream_message_reconnect(request, message_uuid: str, offset: int = 0): + api_key = _get_api_key( + request, + select_related=['user__host_account', 'user__business_account'], + check_usage_limit=False, + ) + user = api_key.user + if user.account_type not in ('business_host', 'regular', 'business_admin'): + raise HttpError(403, _('API key is not available for this account type')) + + store = PublicSSEStoreService(user_uuid=user.pk, message_uuid=message_uuid) + 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', 'X-Accel-Buffering': 'no'}, + ) + + +@router.post('text/{model_slug}/stream', tags=['public/text']) +def public_stream_message(request, model_slug: str, body: MessageInSchema): + api_key = _get_api_key( + request, + select_related=[ + 'user__host_account', + 'user__business_account', + 'user__payment_plan', + 'user__payment_plan__plan', + 'user__business_account__parent_company', + 'user__business_account__parent_company__user__payment_plan', + 'user__business_account__parent_company__user__payment_plan__plan', + ], + prefetch_related=[ + 'user__payment_plan__plan__features', + 'user__business_account__parent_company__user__payment_plan__plan__features', + ], + ) + user = api_key.user + if user.account_type not in ('business_host', 'regular', 'business_admin'): + raise HttpError(403, _('API key is not available for this account type')) + balance = user.balance + + api_store, created = APIStore.objects.get_or_create(user=user) + + selector = NeuronModelSelector(user) + model = selector.get_model_by_slug(slug=model_slug) + if model.blocked: + raise HttpError(403, _('Model is blocked by outdating or temporary block, please retry later')) + if not model.streaming: + raise HttpError(501, _('Stream not supported for this model')) + + if not body.content: + raise HttpError(400, _('The request must not be empty')) + + data = body.model_dump(include={'content', 'file', 'info'}, exclude_unset=True) + input_message = Message.objects.create( + content_object=api_store, from_model=False, from_public_api=True, **data + ) + store = PublicSSEStoreService(user_uuid=user.pk, message_uuid=input_message.pk) + store.start() + + public_event_stream_task.delay( + start_user_balance=balance, + message_uuid=str(input_message.pk), + user_uuid=str(user.pk), + model_slug=model_slug, + api_key_uuid=str(api_key.pk), + debit_api_key_limit=api_key.token_limit is not None, + ) + + sse_chat_stream = SSEChatStreamService(store) + return StreamingHttpResponse( + sse_chat_stream.event_stream(request=request), + content_type='text/event-stream', + headers={'Cache-Control': 'no-cache', 'X-Accel-Buffering': 'no'}, + ) + + +from tools.public_api.routes.providers.openai import openai_responses_stream + + +@router.post('openai/v1/responses', tags=['openai/responses']) +@router.post('openai/responses', tags=['openai/responses']) +def openai_responses(request, body: dict = Body(...)): + return openai_responses_stream(request, body) @@ -1 +1,3 @@ from .api_key import APIKeyService +from .openai_errors import OpenAIErrorService +from .openai_stream import OpenAIStreamService @@ -0,0 +1,29 @@ +from functools import wraps + +from django.http import JsonResponse +from ninja.errors import HttpError + + +class OpenAIErrorService: + TYPES = { + 400: 'invalid_request_error', 401: 'authentication_error', 403: 'permission_error', + 404: 'invalid_request_error', 409: 'invalid_request_error', 501: 'api_error', + } + + @classmethod + def response(cls, status: int, message: str) -> JsonResponse: + return JsonResponse( + {'error': {'message': str(message), 'type': cls.TYPES.get(status, 'api_error'), 'param': None, 'code': None}}, + status=status, + ) + + @classmethod + def view(cls, fn): + @wraps(fn) + def wrapper(*args, **kwargs): + try: + return fn(*args, **kwargs) + except HttpError as exc: + return cls.response(exc.status_code, exc.message) + + return wrapper @@ -0,0 +1,164 @@ +import time +from typing import Iterator +from uuid import UUID + +import orjson + +from tools.chats.services.sse_chat_stream import SSEChatStreamService +from tools.chats.services.sse_store import PublicSSEStoreService + + +class OpenAIStreamService: + SSE_HEADERS = {'Cache-Control': 'no-cache', 'X-Accel-Buffering': 'no'} + SETUP_N = 3 + _IDX = {'output_index': 0, 'content_index': 0} + + @classmethod + def public_offset(cls, starting_after: int) -> int: + return 0 if starting_after < cls.SETUP_N - 1 else starting_after - cls.SETUP_N + 2 + + @classmethod + def translate( + cls, + public_stream, + model_ref: str, + *, + message_uuid: str | None = None, + starting_after: int = 0, + user_uuid: UUID | None = None, + ) -> Iterator[str]: + min_seq = starting_after or -1 + emit_setup = starting_after < cls.SETUP_N - 1 + meta = cls._meta(message_uuid, model_ref) if message_uuid else None + setup_sent = emit_setup and bool(meta) + tokens: list[str] = [] + last_seq = min_seq + + if setup_sent: + yield from cls._setup(meta, min_seq) + + for chunk in public_stream: + if chunk == SSEChatStreamService.HEARTBEAT: + continue + event, data, eid = cls._parse(chunk) + + if event == 'start': + meta = meta or cls._meta(data['message_uuid'], model_ref) + if emit_setup and not setup_sent: + yield from cls._setup(meta, min_seq) + setup_sent = True + continue + if event == 'pending' or not meta: + continue + + ctx = {'item_id': meta['mid'], **cls._IDX} + if event == 'token' and eid and (token := data.get('content', '')): + tokens.append(token) + if (seq := cls.SETUP_N + eid - 2) > min_seq: + yield cls._sse('response.output_text.delta', seq, delta=token, **ctx) + last_seq = seq + elif event == 'done': + text = data.get('content') or ''.join(tokens) + first = cls.SETUP_N + eid - 2 if eid else cls.SETUP_N + yield from cls._close(meta, text, first, min_seq, ctx) + cls._cleanup_public_stream(user_uuid, meta['rid']) + return + elif event == 'error': + if (seq := last_seq + 1) > min_seq: + yield cls._sse( + 'response.failed', + seq, + response={ + 'id': meta['rid'], + 'object': 'response', + 'status': 'failed', + 'error': {'code': 'server_error', 'message': str(data.get('detail', ''))}, + }, + ) + cls._cleanup_public_stream(user_uuid, meta['rid']) + return + + @classmethod + def _cleanup_public_stream(cls, user_uuid: UUID | None, message_uuid: str | None) -> None: + if user_uuid and message_uuid: + PublicSSEStoreService(user_uuid=user_uuid, message_uuid=message_uuid).cleanup() + + @classmethod + def _meta(cls, message_uuid: str, model: str) -> dict: + uid = str(message_uuid) + return {'rid': uid, 'mid': f'msg_{uid.replace("-", "")}', 'ts': int(time.time()), 'model': model} + + @classmethod + def _sse(cls, event_type: str, seq: int, **fields) -> str: + payload = {'type': event_type, 'sequence_number': seq, **fields} + return f'event: {event_type}\ndata: {orjson.dumps(payload).decode()}\n\n' + + @staticmethod + def _parse(chunk: str) -> tuple[str, dict, int | None]: + event, data, eid = '', {}, None + for line in chunk.split('\n'): + if line.startswith('id:'): + eid = int(line[3:].strip()) + elif line.startswith('event:'): + event = line[6:].strip() + elif line.startswith('data:'): + data = orjson.loads(line[5:].strip()) + return event, data, eid + + @classmethod + def _resp(cls, meta: dict, status: str, output: list) -> dict: + return { + 'id': meta['rid'], + 'object': 'response', + 'created_at': meta['ts'], + 'model': meta['model'], + 'status': status, + 'output': output, + 'store': True, + 'text': {'format': {'type': 'text'}}, + } + + @classmethod + def _setup(cls, meta: dict, min_seq: int) -> Iterator[str]: + item = { + 'id': meta['mid'], + 'type': 'message', + 'status': 'in_progress', + 'role': 'assistant', + 'content': [], + } + for seq, (etype, fields) in enumerate( + ( + ('response.created', {'response': cls._resp(meta, 'in_progress', [])}), + ('response.output_item.added', {'item': item, **cls._IDX}), + ( + 'response.content_part.added', + { + 'item_id': meta['mid'], + 'part': {'type': 'output_text', 'text': '', 'annotations': []}, + **cls._IDX, + }, + ), + ) + ): + if seq > min_seq: + yield cls._sse(etype, seq, **fields) + + @classmethod + def _close(cls, meta: dict, text: str, first_seq: int, min_seq: int, ctx: dict) -> Iterator[str]: + part = {'type': 'output_text', 'text': text, 'annotations': []} + item = { + 'id': meta['mid'], + 'type': 'message', + 'status': 'completed', + 'role': 'assistant', + 'content': [part], + } + for i, (etype, fields) in enumerate( + ( + ('response.output_text.done', {'text': text, **ctx}), + ('response.completed', {'response': cls._resp(meta, 'completed', [item])}), + ) + ): + if (seq := first_seq + i) > min_seq: + yield cls._sse(etype, seq, **fields) @@ -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 ]