@@ -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' @@ -483,3 +488,6 @@ UNLEASH_WEBHOOK_SECRET_KEY = env.str('UNLEASH_WEBHOOK_SECRET_KEY', 'defaultsecre # RECURRING SETTINGS MAX_RECURRING_ATTEMPTS = env.int('MAX_RECURRING_ATTEMPTS', 1) + +# SSE STREAMING +FF__STREAMING_ENABLED = env.bool('FF__STREAMING_ENABLED', False) @@ -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()) @@ -1,7 +1,8 @@ +import json from enum import StrEnum import logging import time -from typing import Any, TypeAlias, TypedDict +from typing import Any, Generator, TypeAlias, TypedDict import httpx @@ -69,6 +70,7 @@ class BytedanceVideoTaskResponse(TypedDict, total=False): RunChatResult: TypeAlias = tuple[str, int, int] +RunStreamChatResult: TypeAlias = Generator[str, None, BytedanceUsage] RunImageResult: TypeAlias = list[str] RunVideoResult: TypeAlias = tuple[str, int] BytedanceRunResult: TypeAlias = RunChatResult | RunImageResult | RunVideoResult @@ -247,6 +249,74 @@ class BytedanceModelArkAdapter: raise GenerationException + @classmethod + def _stream_chat( + cls, + model: str, + callback_data: dict[str, Any], + messages: list[dict[str, Any]] | None = None, + include_reasoning: bool = False, + ) -> RunStreamChatResult: + payload = dict( + **callback_data, + model=model, + messages=messages, + stream=True, + stream_options={'include_usage': True}, + ) + for proxy in Proxy.objects.all(): + with httpx.Client( + base_url=cls.BASE_URL, + headers={ + 'Authorization': f'Bearer {settings.BYTEDANCE_MODEL_ARK_API_KEY}', + 'Content-Type': 'application/json', + }, + proxy=f'{proxy.protocol}://{proxy.address}', + timeout=600, + ) as client: + with client.stream( + 'POST', cls.CONTENT_TYPE_TO_ENDPOINT[BytedanceContentType.CHAT], json=payload + ) as resp: + if resp.status_code >= 400: + logger.error( + 'Bytedance request failed status=%s route=%s body=%s', + resp.status_code, + cls.CONTENT_TYPE_TO_ENDPOINT[BytedanceContentType.CHAT], + resp.text, + ) + raise GenerationException + usage: BytedanceUsage = {} + for line in resp.iter_lines(): + if not line: + continue + if isinstance(line, bytes): + line = line.decode('utf-8') + line = line.strip() + if not line.startswith('data:'): + continue + data = line[5:].lstrip() + if data == '[DONE]': + break + try: + data_obj = json.loads(data) + if choices := data_obj.get('choices'): + chunk = choices[0].get('delta', {}).get('content') or '' + if chunk: + yield chunk + if ( + fr := choices[0].get('finish_reason') + ) and fr != BytedanceFinishReason.STOP: + cls._raise_by_error_payload(data_obj, choices) + if raw_usage := data_obj.get('usage'): + usage = { + 'prompt_tokens': int(raw_usage.get('prompt_tokens') or 0), + 'completion_tokens': int(raw_usage.get('completion_tokens') or 0), + } + except json.JSONDecodeError: + continue + return usage + raise GenerationException + @classmethod def _generate_video( cls, @@ -288,7 +358,7 @@ class BytedanceModelArkAdapter: if error_code := data.get('error', {}).get('code', ''): if error_code == 'OutputImageSensitiveContentDetected': raise NSFWDetectedException - + if image_data := data.get('data'): urls = [item.get('url') for item in image_data if isinstance(item, dict) and item.get('url')] if urls: @@ -398,3 +468,21 @@ class BytedanceModelArkAdapter: raise GenerationException time.sleep(cls.VIDEO_POLL_DELAY_SECONDS) raise GenerationException + + @classmethod + def tokenize(cls, model: str, text: str): + for proxy in Proxy.objects.all(): + with httpx.Client( + base_url=cls.BASE_URL, + headers={ + 'Authorization': f'Bearer {settings.BYTEDANCE_MODEL_ARK_API_KEY}', + 'Content-Type': 'application/json', + }, + proxy=f'{proxy.protocol}://{proxy.address}', + timeout=600, + ) as client: + try: + response = client.post('tokenization', json={'model': model, 'text': [text]}).json() + return response['data'][0]['total_tokens'] + except: + return 0 @@ -1,7 +1,7 @@ import json import logging import math -from typing import Any +from typing import Any, Iterator import httpx import tiktoken @@ -35,7 +35,7 @@ class OpenrouterAdapter: @classmethod def run_streaming_api( cls, version: str, messages: list, callback_data: dict, model_name: str - ) -> ModelResponse: + ) -> Iterator[str]: for proxy in Proxy.objects.all(): with httpx.Client( base_url=cls.BASE_URL, @@ -70,8 +70,11 @@ class OpenrouterAdapter: try: data_obj = json.loads(data) - content += data_obj['choices'][0]['delta'].get('content') or '' + chunk = data_obj['choices'][0]['delta'].get('content') or '' reasoning += data_obj['choices'][0]['delta'].get('reasoning') or '' + if chunk: + content += chunk + yield chunk if data_obj.get('usage'): input_tokens = data_obj['usage']['prompt_tokens'] output_tokens = data_obj['usage']['completion_tokens'] @@ -88,7 +91,22 @@ class OpenrouterAdapter: model_name, messages, content + reasoning ) - return ModelResponse(content, input_tokens, output_tokens) + return input_tokens, output_tokens + + @classmethod + def collect_streaming_api( + cls, version: str, messages: list, callback_data: dict, model_name: str + ) -> ModelResponse: + stream = cls.run_streaming_api(version, messages, callback_data, model_name) + content_parts: list[str] = [] + try: + while True: + chunk = next(stream) + if chunk: + content_parts.append(chunk) + except StopIteration as exc: + input_tokens, output_tokens = exc.value + return ModelResponse(''.join(content_parts), input_tokens, output_tokens) @classmethod def _fallback_tokenize(cls, model_name: str, messages: list, content: str) -> tuple[int, int]: @@ -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: ... @@ -1,4 +1,6 @@ import base64 +from typing import Any, Iterator + import filetype import httpx import time @@ -16,20 +18,27 @@ from langchain_core.messages import AIMessage, BaseMessage, HumanMessage, System from messages.models import Message -from ml_model.exceptions import CorruptedFileError, FileExtensionNotSupported, InvalidParameterError, PaidPlanRequiredError +from ml_model.exceptions import ( + CorruptedFileError, + FileExtensionNotSupported, + InvalidParameterError, + PaidPlanRequiredError, +) from ml_model.services import Chatgpt from ml_model.services.EmbeddingService import EmbeddingService from ml_model.services.FileService import FileProcessingService from pathlib import Path +from ml_model.services.base import StreamSimpleService +from ml_model.services.openai_stream_mixin import OpenAIStreamMixin from payments.exceptions.insufficient_balance import InsufficientBalance from payments.selectors.payment_plan_selector import PaymentPlanSelector from poller.models import Proxy -class Chatgpt_5_5(Chatgpt): +class Chatgpt_5_5(Chatgpt, StreamSimpleService, OpenAIStreamMixin): TOKENS_COST = { 'gpt-5.5': { 'input': Decimal('0.0025'), # $5 / 1M tokens @@ -153,196 +162,10 @@ class Chatgpt_5_5(Chatgpt): return price.quantize(Decimal('0.1'), rounding='ROUND_UP') def make(self, input_message: Message, save: bool = True) -> list[Message]: - start_time = time.time() - info = input_message.info.copy() - model_name = 'gpt-5.5' - user_system_prompt = info.pop('system_prompt', '') - file = input_message.file - image = None - embedding_tokens = 0 - predict_embedding_tokens = 0 - chunks = [] - text_chunks = [] - predicted_input_price = 0 - is_free_plan = self.store.user.plan.price <= 0 - if is_free_plan: - info.pop('code_interpreter', None) - info.pop('verbosity', None) - if file: - supported_extensions = ['PDF', 'DOC', 'DOCX', 'XLSX', 'JPG', 'JPEG', 'PNG', 'WEBP'] - file_service = FileProcessingService - file_bytes = input_message.file.read() - kind = filetype.guess(file_bytes[:550]) - if not kind: - if Path(input_message.file.name).suffix[1:].upper() not in supported_extensions: - raise FileExtensionNotSupported(supported_extensions) - raise CorruptedFileError - raw_file_extension = kind.extension - file_extension = file_service.get_file_extension(raw_file_extension, file_bytes) - if file_extension in ('pdf', 'doc', 'docx', 'xlsx'): - if is_free_plan: - raise PaidPlanRequiredError(gettext('File analysis')) - text = file_service.get_file_data(file_extension, file_bytes) - text_chunks = EmbeddingService.split_text_to_chunks(text) - chunks = [HumanMessage(content=chunk_text) for chunk_text in text_chunks] - if (chunks_length := sum(len(chunk) for chunk in text_chunks)) > 20_000: - predict_embedding_tokens = len(text_chunks) * 2020 - predicted_input_price += ( - Decimal((210 + (chunks_length / len(text_chunks) * 10)) / 2.7) - * self.TOKENS_COST[model_name]['input'] - ) - else: - predicted_input_price += ( - Decimal((55 + chunks_length + len(input_message.content)) / 2.2) - * self.TOKENS_COST[model_name]['input'] - ) - elif file_extension in ('jpg', 'jpeg', 'png', 'webp'): - image = file - _, _, image_data = self._get_image_data(file_bytes, file_extension) - else: - raise FileExtensionNotSupported(supported_extensions) - if image and info.get('code_interpreter'): - raise InvalidParameterError( - gettext('The "Use code" option cannot be used together with an attached image.') - ) - chat_history = self.get_chat_history(model_name=model_name) - chat_history.add_message(HumanMessage(content=input_message.content)) - current_user_balance = PaymentPlanSelector(self.store.user).get_current_balance() - is_low_balance = current_user_balance < Decimal('100') - has_full_access = not is_free_plan and not is_low_balance - gpt_5_5_system = SystemMessage(content=self.BASE_SYSTEM) + ctx: dict[str, Any] = {} for proxy in Proxy.objects.all(): - history_messages = list(chat_history.messages) - system = history_messages[0] - messages = [ - { - 'role': 'user' if isinstance(msg, HumanMessage) else 'assistant', - 'content': msg.content, - } - for msg in history_messages[1:] - ] - messages.insert(0, {'role': 'system', 'content': system.content}) - messages.insert(0, {'role': 'system', 'content': gpt_5_5_system.content}) - messages.insert(0, {'role': 'system', 'content': user_system_prompt}) - json_data = { - 'model': model_name, - 'input': messages, - 'instructions': self.FORMATION_INSTRUCTIONS, - 'tools': [], - } - if has_full_access: - json_data['tools'].append( - { - 'type': 'image_generation', - 'size': '1024x1024', - 'quality': 'medium', - 'model': 'gpt-image-1.5', - } - ) - if reasoning := info.get('reasoning'): - reasoning_data = { - 'Средний': 'low', - 'Высокий': 'medium', - } - effort = 'none' if is_free_plan else 'low' if is_low_balance else reasoning_data[reasoning] - json_data['reasoning'] = {'effort': effort, 'summary': 'auto'} - else: - json_data['reasoning'] = {'effort': 'none', 'summary': 'auto'} - web_search = info.get('web_search', 'Выключен') - if info.get('verbosity', False) and web_search in ('Выключен', 'Низкий'): - json_data['text'] = {'verbosity': 'low'} - serper_sources = 0 - if web_search != 'Выключен': - search_context_sizes = { - 'Низкий': 'low', - 'Средний': 'medium', - 'Высокий': 'medium', - 'Сверхвысокий': 'medium', - } - if has_full_access: - json_data['tools'].append( - { - 'type': 'web_search', - 'search_context_size': search_context_sizes[web_search], - 'user_location': {'type': 'approximate', 'country': 'RU'}, - } - ) - info['web_search'] = search_context_sizes[web_search] - predicted_input_price += self.TOKENS_COST[model_name]['web_search'][ - search_context_sizes[web_search] - ] - else: - max_sources = { - 'Низкий': 3, - 'Средний': 5, - 'Высокий': 7, - 'Сверхвысокий': 10, - } - serper_sources = max_sources[web_search] - predicted_input_price += ( - Decimal(serper_sources * 250) - / Decimal(2.7) - * self.TOKENS_COST[model_name]['input'] - ) - if info.get('code_interpreter'): - json_data['tools'].append({'type': 'code_interpreter', 'container': {'type': 'auto'}}) - messages[-1]['content'] += ' the python tool ' - predicted_input_price += self.TOKENS_COST[model_name]['code_interpreter'] - if image: - messages[-1]['content'] = [ - {'type': 'input_text', 'text': input_message.content}, - {'type': 'input_image', 'image_url': image_data['image_url']['url']}, - ] - predicted_input_tokens = self._count_responses_input_tokens(proxy, json_data) - predicted_input_price += ( - predicted_input_tokens * self.TOKENS_COST[model_name]['input'] - + predict_embedding_tokens - * self.TOOLS_TOKEN_COSTS[self.EMBEDDING_MODEL_FOR_BILLING]['output'] - ) - if has_full_access: - predicted_input_price += self.TOKENS_COST[model_name]['generated_image'] - max_output_tokens = max( - int((current_user_balance - predicted_input_price - Decimal('0.1')) / self.TOKENS_COST[model_name]['output']), - 0, - ) - min_response_tokens = 300 if not is_free_plan else 150 - if max_output_tokens < min_response_tokens: - cost = predicted_input_price + min_response_tokens * self.TOKENS_COST[model_name]['output'] - raise InsufficientBalance(current_user_balance, cost) - json_data['max_output_tokens'] = max_output_tokens - if file and not image: - if sum([len(chunk.content) for chunk in chunks]) > 20_000: - document_name = ( - chunks[0].content.partition(f':{chr(10)}')[2].split(f'{chr(10)}')[0][:100] - ) - embedding_tokens, file_data = EmbeddingService.get_large_file_data( - self.store.messages.first().pk, - text_chunks, - proxy, - input_message.content, - model='text-embedding-3-small', - index_name='ml_model-index-1536', - ) - messages[-1]['content'] = EmbeddingService.make_embeddings_prompt( - document_name=document_name, - section_texts=file_data, - question=input_message.content, - ) - else: - messages[-1]['content'] = ( - 'Используй системный промпт. Содержание файла: ' - f'{"".join(text_chunks)}. Вопрос: {input_message.content}' - ) - if serper_sources: - serp = self.run_serper(input_message.content) - messages.append( - { - 'role': 'user', - 'content': self.serper_to_openai_context(serp, max_sources=serper_sources), - } - ) - json_data['input'] = messages - info.pop('web_search', None) + json_data, predicted_input_tokens = self._build_payload(proxy, input_message, ctx) + model_name = ctx['model_name'] self.logger.info( f'Predicted input tokens (responses/input_tokens) для {model_name} - {predicted_input_tokens}' ) @@ -358,27 +181,308 @@ class Chatgpt_5_5(Chatgpt): response.content = gettext_lazy('Image is ready') self.logger.info(f'Input количество токенов для {model_name} - {input_tokens}') self.logger.info(f'Output количество токенов для {model_name} - {output_tokens}') - self.logger.info(f'Embedding количество токенов для {model_name} - {embedding_tokens}') + self.logger.info(f'Embedding количество токенов для {model_name} - {ctx["embedding_tokens"]}') if generated_image: self.logger.info( f'Фиксированная цена за генерацию картинки - ' f'{self.TOKENS_COST[model_name]["generated_image"]}' ) self.logger.info( - f'Общее количество токенов для {model_name} - {input_tokens + output_tokens + embedding_tokens}' + f'Общее количество токенов для {model_name} - ' + f'{input_tokens + output_tokens + ctx["embedding_tokens"]}' ) - process_time = timedelta(seconds=time.time() - start_time) + process_time = timedelta(seconds=time.time() - ctx['start_time']) self.handle_invoice( input_message.content_object.model, input_tokens, output_tokens, model_name, - info, - embedding_tokens, + ctx['info'], + ctx['embedding_tokens'], generated_image, ) - msgs = self.save_results([response], process_time, generated_image, save) - return msgs + return self.save_results([response], process_time, generated_image, save) + + def make_stream(self, input_message: Message, save: bool = True) -> Iterator[str]: + ctx: dict[str, Any] = {} + content_parts: list[str] = [] + input_tokens = output_tokens = 0 + result = '' + + try: + for proxy in Proxy.objects.all(): + json_data, predicted_input_tokens = self._build_payload(proxy, input_message, ctx) + model_name = ctx['model_name'] + self.logger.info( + f'Predicted input tokens (responses/input_tokens) для {model_name} - {predicted_input_tokens}' + ) + stream = self.run_stream(json_data, proxy, model_name) + try: + while True: + chunk = next(stream) + if chunk: + content_parts.append(chunk) + yield chunk + except StopIteration as exc: + input_tokens, output_tokens = exc.value or (0, 0) + break + finally: + if content_parts and ctx: + result = ''.join(content_parts).replace('\\n', '\n') + response = AIMessage(content=result) + if output_tokens == 0: + output_tokens = self.count_text_tokens([response]) + model_name = ctx['model_name'] + self.logger.info(f'Input количество токенов для {model_name} - {input_tokens}') + self.logger.info(f'Output количество токенов для {model_name} - {output_tokens}') + self.logger.info( + f'Embedding количество токенов для {model_name} - {ctx["embedding_tokens"]}' + ) + self.logger.info( + f'Общее количество токенов для {model_name} - ' + f'{input_tokens + output_tokens + ctx["embedding_tokens"]}' + ) + process_time = timedelta(seconds=time.time() - ctx['start_time']) + self.handle_invoice( + input_message.content_object.model, + input_tokens, + output_tokens, + model_name, + ctx['info'], + ctx['embedding_tokens'], + None, + ) + self.save_results([response], process_time, None, save) + if result: + return result + + def _build_payload( + self, + proxy: Proxy, + input_message: Message, + ctx: dict[str, Any], + ) -> tuple[dict[str, Any], int]: + if not ctx: + model_name = 'gpt-5.5' + info = input_message.info.copy() + user_system_prompt = info.pop('system_prompt', '') + file = input_message.file + image = None + image_data = None + predict_embedding_tokens = 0 + predicted_input_price = Decimal(0) + chunks: list[HumanMessage] = [] + text_chunks: list[str] = [] + is_free_plan = self.store.user.plan.price <= 0 + if is_free_plan: + info.pop('code_interpreter', None) + info.pop('verbosity', None) + if file: + supported_extensions = ['PDF', 'DOC', 'DOCX', 'XLSX', 'JPG', 'JPEG', 'PNG', 'WEBP'] + file_service = FileProcessingService + file_bytes = input_message.file.read() + kind = filetype.guess(file_bytes[:550]) + if not kind: + if Path(input_message.file.name).suffix[1:].upper() not in supported_extensions: + raise FileExtensionNotSupported(supported_extensions) + raise CorruptedFileError + raw_file_extension = kind.extension + file_extension = file_service.get_file_extension(raw_file_extension, file_bytes) + if file_extension in ('pdf', 'doc', 'docx', 'xlsx'): + if is_free_plan: + raise PaidPlanRequiredError(gettext('File analysis')) + text = file_service.get_file_data(file_extension, file_bytes) + text_chunks = EmbeddingService.split_text_to_chunks(text) + chunks = [HumanMessage(content=chunk_text) for chunk_text in text_chunks] + if (chunks_length := sum(len(chunk) for chunk in text_chunks)) > 20_000: + predict_embedding_tokens = len(text_chunks) * 2020 + predicted_input_price += ( + Decimal((210 + (chunks_length / len(text_chunks) * 10)) / 2.7) + * self.TOKENS_COST[model_name]['input'] + ) + else: + predicted_input_price += ( + Decimal((55 + chunks_length + len(input_message.content)) / 2.2) + * self.TOKENS_COST[model_name]['input'] + ) + elif file_extension in ('jpg', 'jpeg', 'png', 'webp'): + image = file + _, _, image_data = self._get_image_data(file_bytes, file_extension) + else: + raise FileExtensionNotSupported(supported_extensions) + if image and info.get('code_interpreter'): + raise InvalidParameterError( + gettext('The "Use code" option cannot be used together with an attached image.') + ) + chat_history = self.get_chat_history(model_name=model_name) + chat_history.add_message(HumanMessage(content=input_message.content)) + current_user_balance = PaymentPlanSelector(self.store.user).get_current_balance() + is_low_balance = current_user_balance < Decimal('100') + ctx.update( + { + 'start_time': time.time(), + 'info': info, + 'model_name': model_name, + 'user_system_prompt': user_system_prompt, + 'file': file, + 'image': image, + 'image_data': image_data, + 'embedding_tokens': 0, + 'predict_embedding_tokens': predict_embedding_tokens, + 'predicted_input_price': predicted_input_price, + 'chunks': chunks, + 'text_chunks': text_chunks, + 'chat_history': chat_history, + 'current_user_balance': current_user_balance, + 'is_free_plan': is_free_plan, + 'is_low_balance': is_low_balance, + 'has_full_access': not is_free_plan and not is_low_balance, + 'gpt_5_5_system': SystemMessage(content=self.BASE_SYSTEM), + } + ) + + model_name = ctx['model_name'] + info = ctx['info'] + history_messages = list(ctx['chat_history'].messages) + system = history_messages[0] + messages = [ + { + 'role': 'user' if isinstance(msg, HumanMessage) else 'assistant', + 'content': msg.content, + } + for msg in history_messages[1:] + ] + messages.insert(0, {'role': 'system', 'content': system.content}) + messages.insert(0, {'role': 'system', 'content': ctx['gpt_5_5_system'].content}) + messages.insert(0, {'role': 'system', 'content': ctx['user_system_prompt']}) + json_data = { + 'model': model_name, + 'input': messages, + 'instructions': self.FORMATION_INSTRUCTIONS, + 'tools': [], + } + if ctx['has_full_access']: + json_data['tools'].append( + { + 'type': 'image_generation', + 'size': '1024x1024', + 'quality': 'medium', + 'model': 'gpt-image-1.5', + } + ) + if reasoning := info.get('reasoning'): + reasoning_data = { + 'Средний': 'low', + 'Высокий': 'medium', + } + effort = ( + 'none' + if ctx['is_free_plan'] + else 'low' + if ctx['is_low_balance'] + else reasoning_data[reasoning] + ) + json_data['reasoning'] = {'effort': effort, 'summary': 'auto'} + else: + json_data['reasoning'] = {'effort': 'none', 'summary': 'auto'} + web_search = info.get('web_search', 'Выключен') + if info.get('verbosity', False) and web_search in ('Выключен', 'Низкий'): + json_data['text'] = {'verbosity': 'low'} + serper_sources = 0 + predicted_input_price = ctx['predicted_input_price'] + if web_search != 'Выключен': + search_context_sizes = { + 'Низкий': 'low', + 'Средний': 'medium', + 'Высокий': 'medium', + 'Сверхвысокий': 'medium', + } + if ctx['has_full_access']: + json_data['tools'].append( + { + 'type': 'web_search', + 'search_context_size': search_context_sizes[web_search], + 'user_location': {'type': 'approximate', 'country': 'RU'}, + } + ) + info['web_search'] = search_context_sizes[web_search] + predicted_input_price += self.TOKENS_COST[model_name]['web_search'][ + search_context_sizes[web_search] + ] + else: + max_sources = { + 'Низкий': 3, + 'Средний': 5, + 'Высокий': 7, + 'Сверхвысокий': 10, + } + serper_sources = max_sources[web_search] + predicted_input_price += ( + Decimal(serper_sources * 250) / Decimal(2.7) * self.TOKENS_COST[model_name]['input'] + ) + if info.get('code_interpreter'): + json_data['tools'].append({'type': 'code_interpreter', 'container': {'type': 'auto'}}) + messages[-1]['content'] += ' the python tool ' + predicted_input_price += self.TOKENS_COST[model_name]['code_interpreter'] + if ctx['image']: + messages[-1]['content'] = [ + {'type': 'input_text', 'text': input_message.content}, + {'type': 'input_image', 'image_url': ctx['image_data']['image_url']['url']}, + ] + predicted_input_tokens = self._count_responses_input_tokens(proxy, json_data) + predicted_input_price += ( + predicted_input_tokens * self.TOKENS_COST[model_name]['input'] + + ctx['predict_embedding_tokens'] + * self.TOOLS_TOKEN_COSTS[self.EMBEDDING_MODEL_FOR_BILLING]['output'] + ) + if ctx['has_full_access']: + predicted_input_price += self.TOKENS_COST[model_name]['generated_image'] + max_output_tokens = max( + int( + (ctx['current_user_balance'] - predicted_input_price - Decimal('0.1')) + / self.TOKENS_COST[model_name]['output'] + ), + 0, + ) + min_response_tokens = 300 if not ctx['is_free_plan'] else 150 + if max_output_tokens < min_response_tokens: + cost = predicted_input_price + min_response_tokens * self.TOKENS_COST[model_name]['output'] + raise InsufficientBalance(ctx['current_user_balance'], cost) + json_data['max_output_tokens'] = max_output_tokens + if ctx['file'] and not ctx['image']: + if sum(len(chunk.content) for chunk in ctx['chunks']) > 20_000: + document_name = ( + ctx['chunks'][0].content.partition(f':{chr(10)}')[2].split(f'{chr(10)}')[0][:100] + ) + ctx['embedding_tokens'], file_data = EmbeddingService.get_large_file_data( + self.store.messages.first().pk, + ctx['text_chunks'], + proxy, + input_message.content, + model='text-embedding-3-small', + index_name='ml_model-index-1536', + ) + messages[-1]['content'] = EmbeddingService.make_embeddings_prompt( + document_name=document_name, + section_texts=file_data, + question=input_message.content, + ) + else: + messages[-1]['content'] = ( + 'Используй системный промпт. Содержание файла: ' + f'{"".join(ctx["text_chunks"])}. Вопрос: {input_message.content}' + ) + if serper_sources: + serp = self.run_serper(input_message.content) + messages.append( + { + 'role': 'user', + 'content': self.serper_to_openai_context(serp, max_sources=serper_sources), + } + ) + json_data['input'] = messages + info.pop('web_search', None) + return json_data, predicted_input_tokens def _count_responses_input_tokens(self, proxy: Proxy, json_data: dict) -> int: payload = {k: v for k, v in json_data.items() if k not in ('stream', 'max_output_tokens')} @@ -3,16 +3,18 @@ import time from datetime import timedelta from decimal import Decimal from io import BytesIO +from pathlib import Path from typing import Any, Iterator import filetype from PIL import Image from messages.models import Message -from ml_model.exceptions import FileExtensionNotSupported, FileUploadUnsupported, ModelVersionNotAvailable +from ml_model.adapters.openrouter import OpenrouterAdapter +from ml_model.exceptions import CorruptedFileError, FileExtensionNotSupported, ModelVersionNotAvailable from ml_model.services.EmbeddingService import EmbeddingService from ml_model.services.FileService import FileProcessingService -from ml_model.services.base import SimpleService +from ml_model.services.base import StreamSimpleService from ml_model.tasks import openrouter_run from payments.exceptions.insufficient_balance import InsufficientBalance from payments.selectors.payment_plan_selector import PaymentPlanSelector @@ -21,8 +23,7 @@ from tools.chats.models import Chat from tools.copywrite.models import Copywrite from tools.public_api.models import APIStore - -class Claude(SimpleService): +class Claude(StreamSimpleService): """ Claude Service contains abstract method make, which makes a generation @@ -71,10 +72,12 @@ class Claude(SimpleService): TOOLS_TOKEN_COSTS = {'text-embedding-3-small': {'output': Decimal('0.00001')}} + SUPPORTED_EXTENSIONS = ['PDF', 'DOC', 'DOCX', 'XLSX', 'JPG', 'JPEG', 'PNG', 'WEBP'] + def calculate_price( self, version: str, input_tokens: int, output_tokens: int, embedding_tokens: int ) -> Decimal: - price_map = self.TOKENS_COST[version.split('/')[1]] + price_map = self.TOKENS_COST[version] price = ( input_tokens * price_map['input'] / 1_000_000 + output_tokens * price_map['output'] / 1_000_000 ) @@ -83,7 +86,7 @@ class Claude(SimpleService): price += Decimal('2') return price.quantize(Decimal('0.1'), rounding='ROUND_UP') - def save_results(self, content: Iterator[Any], t: timedelta, save: bool = True) -> list[Message]: + def save_results(self, content: str, t: timedelta, save: bool = True) -> list[Message]: msgs = [ Message( content=content, @@ -97,102 +100,61 @@ class Claude(SimpleService): def make(self, input_message: Message, save: bool = True) -> list[Message]: start_time = time.time() - version_slug = input_message.info.pop('version', None) + version_slug = input_message.info.get('version') if version_slug is None or version_slug not in self.TOKENS_COST: raise ModelVersionNotAvailable(version_slug, self.TOKENS_COST) - version = f'anthropic/{version_slug}' - system_prompt = input_message.info.pop('system_prompt', '') - callback_data = {'provider': {'order': ['anthropic']}, **input_message.info, 'tools': []} - messages = [ - {'role': 'system', 'content': system_prompt}, - *self.get_chat_history(), - {'role': 'user', 'content': input_message.content}, - ] - messages.insert(0, {'role': 'system', 'content': self.HISTORY_CONTEXT_SYSTEM_PROMPT}) - if version_slug == 'claude-fable-5': - current_user_balance = PaymentPlanSelector(self.store.user).get_current_balance() - if current_user_balance < (cost := Decimal('100')) and input_message.file: - raise InsufficientBalance(current_user_balance, cost) - messages.insert(0, {'role': 'system', 'content': self.FABLE_SYSTEM_PROMPT}) - if reasoning := input_message.info.get('reasoning'): - reasoning_data = { - 'Средний': 'low', - 'Высокий': 'medium', - } - callback_data['reasoning'] = {'effort': reasoning_data[reasoning]} - callback_data['tools'].append( - { - 'type': 'openrouter:web_search', - 'parameters': { - 'engine': 'parallel', - 'max_results': 1, - 'max_total_results': 3, - 'search_context_size': 'low', - }, - } - ) - file = input_message.file - embedding_tokens = 0 - if file: - file_service = FileProcessingService - try: - file_bytes = input_message.file.read() - kind = filetype.guess(file_bytes[:550]) - raw_file_extension = kind.extension - file_extension = file_service.get_file_extension(raw_file_extension, file_bytes) - if file_extension in ('pdf', 'doc', 'docx', 'xlsx'): - text = file_service.get_file_data(file_extension, file_bytes) - chunks = EmbeddingService.split_text_to_chunks(text) - if len(text) > 20_000: - for proxy in Proxy.objects.all(): - document_name = ( - chunks[0].partition(f':{chr(10)}')[2].split(f'{chr(10)}')[0][:100] - ) - embedding_tokens, file_data = EmbeddingService.get_large_file_data( - self.store.messages.first().pk, - chunks, - proxy, - input_message.content, - model='text-embedding-3-small', - index_name='ml_model-index-1536', - ) - messages[-1]['content'] = EmbeddingService.make_embeddings_prompt( - document_name=document_name, - section_texts=file_data, - question=input_message.content, - ) - else: - messages[-1]['content'] = ( - f'Используй системный промпт. Содержание файла: ' - f'{chunks}. Вопрос: {input_message.content}' - ) - else: - kind = filetype.guess(file_bytes[:20]) - mime = kind.mime if kind else 'application/octet-stream' - format = 'jpeg' if kind.extension == 'jpg' else kind.extension - with Image.open(file) as normalized_image: - with BytesIO() as buf: - normalized_image.save(buf, format=format) - image_url = ( - f'data:{mime};base64,{base64.b64encode(buf.getvalue()).decode("utf-8")}' - ) - messages[-1]['content'] = [ - {'type': 'text', 'text': input_message.content}, - {'type': 'image_url', 'image_url': {'url': image_url}}, - ] - except Exception: - raise FileExtensionNotSupported(['PDF', 'DOC', 'DOCX', 'XLSX', 'JPG', 'JPEG', 'PNG']) - result = openrouter_run(f'{version}', messages, callback_data, 'Claude') + model_slug = f'anthropic/{version_slug}' + callback_data = self._build_callback_data(input_message, version_slug) + messages, embedding_tokens = self._prepare_messages(input_message, version_slug) + + result = openrouter_run(model_slug, messages, callback_data, 'Claude') + process_time = timedelta(seconds=(time.time() - start_time)) self.handle_invoice( input_message.content_object.model, - version=version, + version=version_slug, input_tokens=result[1], output_tokens=result[2], embedding_tokens=embedding_tokens, ) - msgs = self.save_results(result[0], process_time) - return msgs + return self.save_results(result[0], process_time, save) + + def make_stream(self, input_message: Message, save: bool = True) -> Iterator[str]: + start_time = time.time() + version_slug = input_message.info.get('version') + if version_slug is None or version_slug not in self.TOKENS_COST: + raise ModelVersionNotAvailable(version_slug, self.TOKENS_COST) + model_slug = f'anthropic/{version_slug}' + callback_data = self._build_callback_data(input_message, version_slug) + messages, embedding_tokens = self._prepare_messages(input_message, version_slug) + content_parts: list[str] = [] + input_tokens = output_tokens = 0 + result = '' + + try: + stream = OpenrouterAdapter.run_streaming_api(model_slug, messages, callback_data, 'Claude') + try: + while True: + chunk = next(stream) + if chunk: + content_parts.append(chunk) + yield chunk + except StopIteration as exc: + input_tokens, output_tokens = exc.value + finally: + if content_parts: + result = ''.join(content_parts) + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice( + input_message.content_object.model, + version=version_slug, + input_tokens=input_tokens, + output_tokens=output_tokens, + embedding_tokens=embedding_tokens, + ) + self.save_results(result, process_time, save) + if result: + return result def get_chat_history( self, message_limit: int = 10, max_character_limit: int = 1500 @@ -217,6 +179,8 @@ class Claude(SimpleService): ).order_by('-created_at')[:message_limit] ) ) + else: + air_messages = [] memory = [] for msg in air_messages: content = msg.content or '' @@ -229,3 +193,95 @@ class Claude(SimpleService): character_length -= len(memory.pop(0)['content']) return memory + + def _build_callback_data(self, input_message: Message, version_slug: str) -> dict[str, Any]: + callback_data = {'provider': {'order': ['anthropic']}, **input_message.info, 'tools': []} + if version_slug == 'claude-fable-5': + if reasoning := input_message.info.get('reasoning'): + reasoning_data = { + 'Средний': 'low', + 'Высокий': 'medium', + } + callback_data['reasoning'] = {'effort': reasoning_data[reasoning]} + callback_data['tools'].append( + { + 'type': 'openrouter:web_search', + 'parameters': { + 'engine': 'parallel', + 'max_results': 1, + 'max_total_results': 3, + 'search_context_size': 'low', + }, + } + ) + return callback_data + + def _prepare_messages( + self, input_message: Message, version_slug: str + ) -> tuple[list[dict[str, str | list]], int]: + if version_slug == 'claude-fable-5' and input_message.file: + current_user_balance = PaymentPlanSelector(self.store.user).get_current_balance() + if current_user_balance < (cost := Decimal('100')): + raise InsufficientBalance(current_user_balance, cost) + + system_prompt = input_message.info.get('system_prompt', '') + messages = [ + {'role': 'system', 'content': self.HISTORY_CONTEXT_SYSTEM_PROMPT}, + {'role': 'system', 'content': system_prompt}, + *self.get_chat_history(), + {'role': 'user', 'content': input_message.content}, + ] + if version_slug == 'claude-fable-5': + messages.insert(1, {'role': 'system', 'content': self.FABLE_SYSTEM_PROMPT}) + + embedding_tokens = 0 + if input_message.file: + file_service = FileProcessingService + file_bytes = input_message.file.read() + kind = filetype.guess(file_bytes[:550]) + if not kind: + if Path(input_message.file.name).suffix[1:].upper() not in self.SUPPORTED_EXTENSIONS: + raise FileExtensionNotSupported(self.SUPPORTED_EXTENSIONS) + raise CorruptedFileError + raw_file_extension = kind.extension + file_extension = file_service.get_file_extension(raw_file_extension, file_bytes) + if file_extension in ('pdf', 'doc', 'docx', 'xlsx'): + text = file_service.get_file_data(file_extension, file_bytes) + chunks = EmbeddingService.split_text_to_chunks(text) + if len(text) > 20_000: + for proxy in Proxy.objects.all(): + document_name = chunks[0].partition(f':{chr(10)}')[2].split(f'{chr(10)}')[0][:100] + embedding_tokens, file_data = EmbeddingService.get_large_file_data( + self.store.messages.first().pk, + chunks, + proxy, + input_message.content, + model='text-embedding-3-small', + index_name='ml_model-index-1536', + ) + messages[-1]['content'] = EmbeddingService.make_embeddings_prompt( + document_name=document_name, + section_texts=file_data, + question=input_message.content, + ) + else: + messages[-1]['content'] = ( + f'Используй системный промпт. Содержание файла: ' + f'{chunks}. Вопрос: {input_message.content}' + ) + elif file_extension in ('jpg', 'jpeg', 'png', 'webp'): + kind = filetype.guess(file_bytes[:20]) + mime = kind.mime if kind else 'application/octet-stream' + format = 'jpeg' if kind.extension == 'jpg' else kind.extension + with Image.open(input_message.file) as normalized_image: + with BytesIO() as buf: + normalized_image.save(buf, format=format) + image_url = f'data:{mime};base64,{base64.b64encode(buf.getvalue()).decode("utf-8")}' + messages[-1]['content'] = [ + {'type': 'text', 'text': input_message.content}, + {'type': 'image_url', 'image_url': {'url': image_url}}, + ] + else: + raise FileExtensionNotSupported(self.SUPPORTED_EXTENSIONS) + + return messages, embedding_tokens @@ -4,16 +4,17 @@ from datetime import timedelta from decimal import Decimal from io import BytesIO from pathlib import Path +from typing import Iterator import filetype from PIL import Image from messages.models import Message -from ml_model.exceptions import CorruptedFileError, FileExtensionNotSupported +from ml_model.adapters.openrouter import OpenrouterAdapter +from ml_model.exceptions import CorruptedFileError, FileExtensionNotSupported, ModelVersionNotAvailable from ml_model.services.EmbeddingService import EmbeddingService from ml_model.services.FileService import FileProcessingService -from ml_model.exceptions import ModelVersionNotAvailable -from ml_model.services.base import SimpleService +from ml_model.services.base import StreamSimpleService from ml_model.tasks import openrouter_run from poller.models import Proxy from tools.chats.models import Chat @@ -21,7 +22,8 @@ from tools.copywrite.models import Copywrite from tools.public_api.models import APIStore -class Gemini_3_1(SimpleService): + +class Gemini_3_1(StreamSimpleService): TOKENS_COST = { 'gemini-3.1-pro-preview': { 'input': Decimal('600'), @@ -37,6 +39,8 @@ class Gemini_3_1(SimpleService): TOOLS_TOKEN_COSTS = {'text-embedding-3-small': {'output': Decimal('0.00001')}} + SUPPORTED_EXTENSIONS = ['PDF', 'DOC', 'DOCX', 'XLSX', 'JPG', 'JPEG', 'PNG', 'WEBP'] + def calculate_price(self, version: str, input_tokens: int, output_tokens: int, embedding_tokens: int) -> Decimal: price_map = self.TOKENS_COST[version] if input_tokens >= 200_000 or output_tokens >= 200_000: @@ -67,6 +71,31 @@ class Gemini_3_1(SimpleService): return msgs def make(self, input_message: Message, save: bool = True) -> list[Message]: + start_time = time.time() + version_slug = input_message.info.get('version') + if version_slug is None or version_slug not in self.TOKENS_COST: + raise ModelVersionNotAvailable(version_slug, self.TOKENS_COST) + model_slug = f'google/{version_slug}:online' + callback_data = { + 'provider': {'order': ['Google AI Studio']}, + **input_message.info, + } + messages, embedding_tokens = self._prepare_messages(input_message) + + result = openrouter_run(model_slug, messages, callback_data, 'Gemini') + + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice( + input_message.content_object.model, + version=version_slug, + input_tokens=result[1], + output_tokens=result[2], + embedding_tokens=embedding_tokens, + ) + return self.save_results(result[0], process_time, save) + + def make_stream(self, input_message: Message, save: bool = True) -> Iterator[str]: + start_time = time.time() version_slug = input_message.info.get('version') if version_slug is None or version_slug not in self.TOKENS_COST: raise ModelVersionNotAvailable(version_slug, self.TOKENS_COST) @@ -75,6 +104,76 @@ class Gemini_3_1(SimpleService): 'provider': {'order': ['Google AI Studio']}, **input_message.info, } + messages, embedding_tokens = self._prepare_messages(input_message) + content_parts: list[str] = [] + input_tokens = output_tokens = 0 + result = '' + + try: + stream = OpenrouterAdapter.run_streaming_api(model_slug, messages, callback_data, 'Gemini') + try: + while True: + chunk = next(stream) + if chunk: + content_parts.append(chunk) + yield chunk + except StopIteration as exc: + input_tokens, output_tokens = exc.value + finally: + if content_parts: + result = ''.join(content_parts) + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice( + input_message.content_object.model, + version=version_slug, + input_tokens=input_tokens, + output_tokens=output_tokens, + embedding_tokens=embedding_tokens, + ) + self.save_results(result, process_time, save) + if result: + return result + + def get_chat_history( + self, message_limit: int = 10, max_character_limit: int = 1500 + ) -> list[dict[str, str | list]]: + if isinstance(self.store, Chat): + air_messages = list( + reversed( + Message.objects.filter( + chats_chats_messages=self.store, is_deleted=False, is_sent=True + ).order_by('-created_at')[1 : message_limit + 1] + ) + ) + elif isinstance(self.store, APIStore): + air_messages = [] + elif isinstance(self.store, Copywrite): + air_messages = list( + reversed( + Message.objects.filter( + copywrite_copywrites_messages=self.store, + is_deleted=False, + is_sent=True, + ).order_by('-created_at')[:message_limit] + ) + ) + else: + air_messages = [] + memory = [] + for msg in air_messages: + content = msg.content or '' + if msg.from_model: + memory.append({'role': 'assistant', 'content': content}) + else: + memory.append({'role': 'user', 'content': content}) + character_length = sum(len(content['content']) for content in memory) + while character_length > max_character_limit: + character_length -= len(memory.pop(0)['content']) + return memory + + def _prepare_messages( + self, input_message: Message + ) -> tuple[list[dict[str, str | list]], int]: messages = self.get_chat_history() messages.insert( 0, @@ -89,13 +188,12 @@ class Gemini_3_1(SimpleService): messages.append({'role': 'user', 'content': input_message.content}) embedding_tokens = 0 if input_message.file: - supported_extensions = ['PDF', 'DOC', 'DOCX', 'XLSX', 'JPG', 'JPEG', 'PNG', 'WEBP'] file_service = FileProcessingService file_bytes = input_message.file.read() kind = filetype.guess(file_bytes[:550]) if not kind: - if Path(input_message.file.name).suffix[1:].upper() not in supported_extensions: - raise FileExtensionNotSupported(supported_extensions) + if Path(input_message.file.name).suffix[1:].upper() not in self.SUPPORTED_EXTENSIONS: + raise FileExtensionNotSupported(self.SUPPORTED_EXTENSIONS) raise CorruptedFileError raw_file_extension = kind.extension file_extension = file_service.get_file_extension(raw_file_extension, file_bytes) @@ -136,53 +234,6 @@ class Gemini_3_1(SimpleService): {'type': 'image_url', 'image_url': {'url': image_url}}, ] else: - raise FileExtensionNotSupported(supported_extensions) - start_time = time.time() - result = openrouter_run(model_slug, messages, callback_data, 'Gemini 3.1') - process_time = timedelta(seconds=(time.time() - start_time)) - self.handle_invoice( - input_message.content_object.model, - version=version_slug, - input_tokens=result[1], - output_tokens=result[2], - embedding_tokens=embedding_tokens, - ) - msgs = self.save_results(result[0], process_time) - return msgs + raise FileExtensionNotSupported(self.SUPPORTED_EXTENSIONS) - def get_chat_history( - self, message_limit: int = 10, max_character_limit: int = 1500 - ) -> list[dict[str, str | list]]: - if isinstance(self.store, Chat): - air_messages = list( - reversed( - Message.objects.filter( - chats_chats_messages=self.store, is_deleted=False, is_sent=True - ).order_by('-created_at')[1 : message_limit + 1] - ) - ) - elif isinstance(self.store, APIStore): - air_messages = [] - elif isinstance(self.store, Copywrite): - air_messages = list( - reversed( - Message.objects.filter( - copywrite_copywrites_messages=self.store, - is_deleted=False, - is_sent=True, - ).order_by('-created_at')[:message_limit] - ) - ) - else: - air_messages = [] - memory = [] - for msg in air_messages: - content = msg.content or '' - if msg.from_model: - memory.append({'role': 'assistant', 'content': content}) - else: - memory.append({'role': 'user', 'content': content}) - character_length = sum(len(content['content']) for content in memory) - while character_length > max_character_limit: - character_length -= len(memory.pop(0)['content']) - return memory + return messages, embedding_tokens @@ -6,8 +6,9 @@ from typing import Any, Iterator from messages.models import Message from ml_model.adapters.bytedance_model_ark import BytedanceContentType, BytedanceModelArkAdapter +from ml_model.exceptions import GenerationException from ml_model.services.base import SimpleService -from ml_model.tasks import bytedance_model_ark_run +from ml_model.tasks import bytedance_model_ark_run, stream_bytedance_model_ark_run from tools.chats.models import Chat from tools.copywrite.models import Copywrite from tools.public_api.models import APIStore @@ -31,7 +32,7 @@ class Glm_4_7(SimpleService): ) return price.quantize(Decimal('0.1'), rounding='ROUND_UP') - def save_results(self, content: Iterator[Any], t: timedelta, save: bool = True) -> list[Message]: + def save_results(self, content: str, t: timedelta, save: bool = True) -> list[Message]: msgs = [ Message( content=content, @@ -43,7 +44,7 @@ class Glm_4_7(SimpleService): return Message.objects.bulk_create(msgs) return msgs - def make(self, input_message: Message, save: bool = True) -> list[Message]: + def _prepare_data(self, input_message: Message) -> tuple[str, dict[str, Any], list[dict[str, Any]]]: version = 'glm-4-7-251222' callback_data = { "reasoning_effort": "minimal", @@ -53,7 +54,10 @@ class Glm_4_7(SimpleService): messages.append({'role': 'user', 'content': input_message.content}) if input_message.file: logger.info('GLM file input ignored: model does not support image/video input') + return version, callback_data, messages + def make(self, input_message: Message, save: bool = True) -> list[Message]: + version, callback_data, messages = self._prepare_data(input_message) start_time = time.time() result = bytedance_model_ark_run( @@ -72,7 +76,52 @@ class Glm_4_7(SimpleService): msgs = self.save_results(result[0], process_time, save) return msgs - def get_chat_history(self, message_limit: int = 10, max_character_limit: int = 1500) -> list[dict[str, str | list]]: + def make_stream(self, input_message: Message, save: bool = True) -> Iterator[str]: + version, callback_data, messages = self._prepare_data(input_message) + start_time = time.time() + content_parts: list[str] = [] + input_tokens = output_tokens = 0 + result = '' + + try: + stream = stream_bytedance_model_ark_run( + model=version, + callback_data=callback_data, + messages=messages, + ) + try: + while True: + chunk = next(stream) + if chunk: + content_parts.append(chunk) + yield chunk + except StopIteration as exc: + usage = exc.value or {} + input_tokens = int(usage.get('prompt_tokens') or 0) + output_tokens = int(usage.get('completion_tokens') or 0) + finally: + if content_parts: + result = ''.join(content_parts) + if not (input_tokens + output_tokens): + input_tokens = BytedanceModelArkAdapter.tokenize( + version, ''.join([m['content'] for m in messages]) + ) + output_tokens = BytedanceModelArkAdapter.tokenize(version, result) + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice( + input_message.content_object.model, + version=version, + input_tokens=input_tokens, + output_tokens=output_tokens, + ) + self.save_results(result, process_time, save) + if result: + return result + raise GenerationException + + def get_chat_history( + self, message_limit: int = 10, max_character_limit: int = 1500 + ) -> list[dict[str, str | list]]: if isinstance(self.store, Chat): air_messages = list( reversed( @@ -4,17 +4,17 @@ from datetime import timedelta from decimal import Decimal from io import BytesIO from pathlib import Path -from typing import Any, Iterator +from typing import Iterator import filetype from PIL import Image from messages.models import Message +from ml_model.adapters.openrouter import OpenrouterAdapter from ml_model.exceptions import CorruptedFileError, FileExtensionNotSupported -from ml_model.services.base import SimpleService +from ml_model.services.base import StreamSimpleService from ml_model.services.EmbeddingService import EmbeddingService from ml_model.services.FileService import FileProcessingService -from ml_model.tasks import openrouter_run_streaming from payments.exceptions.insufficient_balance import InsufficientBalance from payments.selectors.payment_plan_selector import PaymentPlanSelector from poller.models import Proxy @@ -23,7 +23,7 @@ from tools.copywrite.models import Copywrite from tools.public_api.models import APIStore -class Grok(SimpleService): +class Grok(StreamSimpleService): TOKENS_COST = { 'grok-4.3': { @@ -36,6 +36,9 @@ class Grok(SimpleService): SUPPORTED_EXTENSIONS = ['PDF', 'DOC', 'DOCX', 'XLSX', 'JPG', 'JPEG', 'PNG', 'WEBP'] + TTFT = 0.5 + TBT = 0.35 + def calculate_price( self, version: str, input_tokens: int, output_tokens: int, embedding_tokens: int ) -> Decimal: @@ -50,22 +53,107 @@ class Grok(SimpleService): price += embedding_tokens * self.TOOLS_TOKEN_COSTS['text-embedding-3-small']['output'] return price.quantize(Decimal('0.1'), rounding='ROUND_UP') - def save_results(self, content: Iterator[Any], t: timedelta, save: bool = True) -> list[Message]: - msgs = [ - Message( - content=content, - content_object=self.store, - elapsed_time=t, - ) - ] + 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: - return Message.objects.bulk_create(msgs) - return msgs + msg.save() + return [msg] def make(self, input_message: Message, save: bool = True) -> list[Message]: start_time = time.time() version = 'x-ai/grok-4.3' callback_data = {**input_message.info} + messages, embedding_tokens = self._prepare_messages(input_message, version) + + result = OpenrouterAdapter.collect_streaming_api(version, messages, callback_data, 'Grok') + + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice( + input_message.content_object.model, + version=version, + input_tokens=result.input_tokens, + output_tokens=result.output_tokens, + embedding_tokens=embedding_tokens, + ) + return self.save_results(result.content, process_time, save) + + def make_stream(self, input_message: Message, save: bool = True) -> Iterator[str]: + start_time = time.time() + version = 'x-ai/grok-4.3' + callback_data = {**input_message.info} + messages, embedding_tokens = self._prepare_messages(input_message, version) + content_parts: list[str] = [] + input_tokens = output_tokens = 0 + result = '' + + try: + stream = OpenrouterAdapter.run_streaming_api(version, messages, callback_data, 'Grok') + try: + while True: + chunk = next(stream) + if chunk: + content_parts.append(chunk) + yield chunk + except StopIteration as exc: + input_tokens, output_tokens = exc.value + finally: + if content_parts: + result = ''.join(content_parts) + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice( + input_message.content_object.model, + version=version, + input_tokens=input_tokens, + output_tokens=output_tokens, + embedding_tokens=embedding_tokens, + ) + self.save_results(result, process_time, save) + if result: + return result + + def get_chat_history( + self, message_limit: int = 10, max_character_limit: int = 1500 + ) -> list[dict[str, str | list]]: + if isinstance(self.store, Chat): + air_messages = list( + reversed( + Message.objects.filter( + chats_chats_messages=self.store, is_deleted=False, is_sent=True + ).order_by('-created_at')[1 : message_limit + 1] + ) + ) + elif isinstance(self.store, APIStore): + air_messages = [] + elif isinstance(self.store, Copywrite): + air_messages = list( + reversed( + Message.objects.filter( + copywrite_copywrites_messages=self.store, + is_deleted=False, + is_sent=True, + ).order_by('-created_at')[:message_limit] + ) + ) + memory = [] + for msg in air_messages: + content = msg.content or '' + if msg.from_model: + memory.append({'role': 'assistant', 'content': content}) + else: + memory.append({'role': 'user', 'content': content}) + character_length = sum(len(content['content']) for content in memory) + while character_length > max_character_limit: + character_length -= len(memory.pop(0)['content']) + + return memory + + def _prepare_messages( + self, input_message: Message, version: str + ) -> tuple[list[dict[str, str | list]], int]: messages = self.get_chat_history() messages.append({'role': 'user', 'content': input_message.content}) embedding_tokens = 0 @@ -130,50 +218,4 @@ class Grok(SimpleService): else: raise FileExtensionNotSupported(self.SUPPORTED_EXTENSIONS) - result = openrouter_run_streaming(version, messages, callback_data, 'Grok') - process_time = timedelta(seconds=(time.time() - start_time)) - self.handle_invoice( - input_message.content_object.model, - version=version, - input_tokens=result.input_tokens, - output_tokens=result.output_tokens, - embedding_tokens=embedding_tokens, - ) - msgs = self.save_results(result.content, process_time) - return msgs - - def get_chat_history( - self, message_limit: int = 10, max_character_limit: int = 1500 - ) -> list[dict[str, str | list]]: - if isinstance(self.store, Chat): - air_messages = list( - reversed( - Message.objects.filter( - chats_chats_messages=self.store, is_deleted=False, is_sent=True - ).order_by('-created_at')[1 : message_limit + 1] - ) - ) - elif isinstance(self.store, APIStore): - air_messages = [] - elif isinstance(self.store, Copywrite): - air_messages = list( - reversed( - Message.objects.filter( - copywrite_copywrites_messages=self.store, - is_deleted=False, - is_sent=True, - ).order_by('-created_at')[:message_limit] - ) - ) - memory = [] - for msg in air_messages: - content = msg.content or '' - if msg.from_model: - memory.append({'role': 'assistant', 'content': content}) - else: - memory.append({'role': 'user', 'content': content}) - character_length = sum(len(content['content']) for content in memory) - while character_length > max_character_limit: - character_length -= len(memory.pop(0)['content']) - - return memory + return messages, embedding_tokens @@ -0,0 +1,111 @@ +import json +from typing import Any + +import httpx + +from django.conf import settings +from poller.models import Proxy + +type OpenAIEvent = dict[str, Any] +type OpenAIUsage = tuple[int, int] +type OpenAIStreamState = dict[str, Any] + + +class OpenAIStreamMixin: + def run_stream(self, json_data: dict[str, Any], proxy: Proxy, model_name: str): + payload = {**json_data, 'stream': True} + state: OpenAIStreamState = {'response_id': '', 'starting_after': -1} + + with httpx.Client( + base_url='https://api.openai.com/v1', + proxy=f'{proxy.protocol}://{proxy.address}', + headers={'Authorization': f'Bearer {settings.OPENAI_API_KEY}'}, + timeout=httpx.Timeout(connect=60, read=600, write=60, pool=30), + ) as client: + try: + input_tokens, output_tokens = yield from self._stream_request( + client, 'POST', json=payload, state=state + ) + except Exception: + if not state['response_id']: + raise + params: dict[str, Any] = {'stream': True} + if state['starting_after'] >= 0: + params['starting_after'] = state['starting_after'] + input_tokens, output_tokens = yield from self._stream_request( + client, + 'GET', + url=f'responses/{state["response_id"]}', + params=params, + state=state, + ) + + return input_tokens, output_tokens + + def _stream_request( + self, + client: httpx.Client, + method: str, + *, + url: str = 'responses', + json: dict[str, Any] | None = None, + params: dict[str, Any] | None = None, + state: OpenAIStreamState, + ): + input_tokens = output_tokens = 0 + + with client.stream(method, url, json=json, params=params) as resp: + resp.raise_for_status() + for line in resp.iter_lines(): + event = self._parse_sse_line(line) + if event is None: + continue + + if (seq := event.get('sequence_number')) is not None: + state['starting_after'] = seq + + event_data = self._get_event_data(event) + if isinstance(event_data, str): + if event_data.startswith('resp_'): + state['response_id'] = event_data + continue + yield event_data + elif isinstance(event_data, tuple): + input_tokens, output_tokens = event_data + break + + return input_tokens, output_tokens + + def _parse_sse_line(self, line: bytes | str) -> OpenAIEvent | None: + if not line: + return None + if isinstance(line, bytes): + line = line.decode('utf-8') + if not line.startswith('data:'): + return None + try: + return json.loads(line[5:].lstrip()) + except ValueError: + return None + + def _get_event_data(self, event: OpenAIEvent) -> str | OpenAIUsage: + event_type = event.get('type') + match event_type: + case 'response.created': + return self._get_response_id(event) + case 'response.output_text.delta': + return self._get_delta(event) + case 'response.completed' | 'response.incomplete' | 'response.failed': + return self._get_usage(event) + case _: + return '' + + def _get_delta(self, event: OpenAIEvent) -> str: + return event.get('delta', '').replace('\n\n\n\n', '\n\n') + + def _get_usage(self, event: OpenAIEvent) -> OpenAIUsage: + usage = event.get('usage') or (event.get('response') or {}).get('usage') or {} + return usage.get('input_tokens', 0), usage.get('output_tokens', 0) + + def _get_response_id(self, event: OpenAIEvent) -> str: + return event.get('response', {}).get('id', '') @@ -1,16 +1,21 @@ +import json import time from datetime import timedelta from decimal import Decimal from pathlib import Path +from typing import Iterator import filetype +import httpx +from backend import settings from messages.models import Message +from ml_model.adapters.openrouter import OpenrouterAdapter from ml_model.exceptions import CorruptedFileError, FileExtensionNotSupported from ml_model.services.EmbeddingService import EmbeddingService from ml_model.services.FileService import FileProcessingService from ml_model.exceptions import ModelVersionNotAvailable -from ml_model.services.base import SimpleService +from ml_model.services.base import StreamSimpleService from ml_model.tasks import openrouter_run from poller.models import Proxy from tools.chats.models import Chat @@ -18,7 +23,7 @@ from tools.copywrite.models import Copywrite from tools.public_api.models import APIStore -class Qwen_3_7(SimpleService): +class Qwen_3_7(StreamSimpleService): COEFFICIENT = Decimal('300.0') TOKENS_COST = { @@ -73,6 +78,61 @@ class Qwen_3_7(SimpleService): def make(self, input_message: Message, save: bool = True) -> list[Message]: start_time = time.time() + version_slug, model_slug, callback_data, messages, embedding_tokens = self._prepare_data( + input_message + ) + result = openrouter_run(model_slug, messages, callback_data, 'Qwen 3.7') + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice( + input_message.content_object.model, + version=version_slug, + cost=result[1], + embedding_tokens=embedding_tokens, + ) + msgs = self.save_results(result[0], process_time) + return msgs + + def make_stream(self, input_message: Message, save: bool = True) -> Iterator[str]: + start_time = time.time() + version_slug, model_slug, callback_data, messages, embedding_tokens = self._prepare_data( + input_message + ) + content_parts: list[str] = [] + cost = 0 + result = '' + + try: + stream = self._run_streaming_api(model_slug, messages, callback_data) + try: + while True: + chunk = next(stream) + if chunk: + content_parts.append(chunk) + yield chunk + except StopIteration as exc: + cost = exc.value + finally: + if content_parts: + result = ''.join(content_parts) + if not cost: + input_tokens, output_tokens = OpenrouterAdapter.count_tokens_fallback( + 'Qwen', messages, result + ) + cost = self._estimate_cost(version_slug, input_tokens, output_tokens) + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice( + input_message.content_object.model, + version=version_slug, + cost=cost, + embedding_tokens=embedding_tokens, + ) + self.save_results(result, process_time, save) + if result: + return result + + def _prepare_data( + self, input_message: Message + ) -> tuple[str, str, dict, list[dict[str, str | list]], int]: version_slug = input_message.info.pop('version', None) if version_slug is None or version_slug not in self.TOKENS_COST: raise ModelVersionNotAvailable(version_slug, self.TOKENS_COST) @@ -133,16 +193,7 @@ class Qwen_3_7(SimpleService): } ) model_slug = f'qwen/{version_slug}' - result = openrouter_run(model_slug, messages, callback_data, 'Qwen 3.7') - process_time = timedelta(seconds=(time.time() - start_time)) - self.handle_invoice( - input_message.content_object.model, - version=version_slug, - cost=result[1], - embedding_tokens=embedding_tokens, - ) - msgs = self.save_results(result[0], process_time) - return msgs + return version_slug, model_slug, callback_data, messages, embedding_tokens def get_chat_history( self, message_limit: int = 10, max_character_limit: int = 1500 @@ -181,3 +232,54 @@ class Qwen_3_7(SimpleService): character_length -= len(memory.pop(0)['content']) return memory + + def _estimate_cost(self, version_slug: str, input_tokens: int, output_tokens: int) -> float: + price_map = self.TOKENS_COST[version_slug] + price = ( + input_tokens * price_map['input'] / 1_000_000 + + output_tokens * price_map['output'] / 1_000_000 + ) + return float(price / self.COEFFICIENT) + + def _run_streaming_api( + self, version: str, messages: list, callback_data: dict + ) -> Iterator[str]: + for proxy in Proxy.objects.all(): + with httpx.Client( + base_url='https://openrouter.ai/api/v1', + headers={'Authorization': f'Bearer {settings.OPENROUTER_API_KEY}'}, + proxy=f'{proxy.protocol}://{proxy.address}', + timeout=600, + ) as client: + with client.stream( + 'POST', + 'chat/completions', + json={ + 'model': version, + 'stream': True, + 'messages': messages, + 'transforms': ['middle-out'], + **callback_data, + }, + ) as resp: + cost = 0 + for line in resp.iter_lines(): + line = line.strip() + if not line or not line.startswith('data: '): + continue + + data = line[6:] + if data == '[DONE]': + break + + try: + data_obj = json.loads(data) + chunk = data_obj['choices'][0]['delta'].get('content') or '' + if chunk: + yield chunk + if data_obj.get('usage'): + cost = data_obj['usage'].get('cost') or 0 + except json.JSONDecodeError: + continue + + return cost @@ -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,5 +1,7 @@ +import importlib from uuid import uuid4 +from django.conf import settings from django.contrib.contenttypes.fields import GenericForeignKey from django.contrib.contenttypes.models import ContentType from django.contrib.postgres.fields import ArrayField @@ -140,6 +142,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 settings.FF__STREAMING_ENABLED and 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): @@ -281,6 +281,18 @@ def bytedance_model_ark_run( ) +def stream_bytedance_model_ark_run( + model: str, + callback_data: dict, + messages: list | None = None, + include_reasoning: bool = False, +): + usage = yield from BytedanceModelArkAdapter._stream_chat( + model, callback_data, messages, include_reasoning + ) + return usage + + @shared_task def drop_redis_vectors(message_uid: str) -> None: redis_client = redis.Redis(host=settings.REDIS_HOST, port=settings.REDIS_PORT, db=0) @@ -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,78 @@ +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, idempotency_key: UUID, user_uuid: UUID): + self.user_uuid = user_uuid + self.idempotency_key = idempotency_key + self.redis_client = self._get_redis_client() + + def _get_cache_key(self) -> str: + return f'sse:tokens:{self.user_uuid}:{self.idempotency_key}' \ 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,96 @@ +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, + idempotency_key: str, + api_key_uuid: str, + debit_api_key_limit: bool, +): + store = PublicSSEStoreService(user_uuid=user_uuid, idempotency_key=idempotency_key) + + 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,148 @@ +import base64 +from io import BytesIO + +import filetype +from django.db.models import Q +from django.http import StreamingHttpResponse +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.chats.services.sse_store import PublicSSEStoreService +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 _stream_context(request, response_id: str | None = None) -> tuple: + from tools.public_api.routes.v1 import _get_api_key + + 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')) + if not (key := (response_id or request.headers.get('Idempotency-Key', '')).strip()): + raise HttpError(400, _('Idempotency-Key header is not provided')) + store = PublicSSEStoreService(user_uuid=user.pk, idempotency_key=key) + return user, store, OpenAIStreamService(store) + + +@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) + _, store, svc = _stream_context(request) + svc.cleanup_finished() + response = public_stream_message(request, model.slug, _parse_body(body)) + meta = svc.init_meta(model_ref, reconnect=False) + response.streaming_content = svc.event_stream(request, meta) + return response + + +@OpenAIErrorService.view +def openai_responses_stream_reconnect(request, response_id: str | None = None, *, stream: bool = True, starting_after: int = 0): + if not stream: + raise HttpError(400, _('Only streaming reconnect is supported')) + user, store, svc = _stream_context(request, response_id) + if not store.exists() and not svc.load_meta(): + raise HttpError(404, _('Stream not found')) + meta = svc.init_meta('', reconnect=True) + if str(meta.get('user_uuid')) != str(user.pk): + raise HttpError(404, _('Stream not found')) + return StreamingHttpResponse( + svc.event_stream(request, meta, reconnect=True, starting_after=starting_after), + content_type='text/event-stream', + headers=OpenAIStreamService.SSE_HEADERS, + ) + + +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,153 @@ +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/{model_slug}/stream', tags=['public/text']) +def public_stream_message_reconnect(request, model_slug: 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')) + + idempotency_key = request.headers.get('Idempotency-Key', '') + if not idempotency_key: + raise HttpError(400, _('Idempotency-Key header is not provided')) + + store = PublicSSEStoreService(user_uuid=user.pk, idempotency_key=idempotency_key) + 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 + + idempotency_key = request.headers.get('Idempotency-Key', '') + if not idempotency_key: + raise HttpError(400, _('Idempotency-Key header is not provided')) + + 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')) + + store = PublicSSEStoreService(user_uuid=user.pk, idempotency_key=idempotency_key) + 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=api_store, from_model=False, from_public_api=True, **data + ) + 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, + idempotency_key=idempotency_key, + 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,180 @@ +import secrets +import time +from typing import Iterator + +import orjson +from django.conf import settings +from django.http import HttpRequest +from django.utils.translation import gettext as _ +from ninja.errors import HttpError + +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} + + def __init__(self, store: PublicSSEStoreService) -> None: + self.store = store + + def cleanup_finished(self) -> None: + if not self.store.exists(): + self._del_meta() + return + key = self.store._get_cache_key() + if any(orjson.loads(r).get('event') in ('done', 'error') for r in self.store.redis_client.lrange(key, 0, -1)): + self.store.cleanup() + self._del_meta() + + def load_meta(self) -> dict | None: + return orjson.loads(r) if (r := self.store.redis_client.get(self._meta_key())) else None + + def init_meta(self, model: str, *, reconnect: bool) -> dict: + if meta := self.load_meta(): + return meta + if reconnect: + raise HttpError(404, _('Stream not found')) + meta = { + 'rid': str(self.store.idempotency_key), + 'mid': f'msg_{secrets.token_hex(24)}', + 'ts': int(time.time()), + 'model': model, + 'user_uuid': str(self.store.user_uuid), + } + self._save_meta(meta) + return meta + + def event_stream( + self, + request: HttpRequest | None, + meta: dict, + *, + reconnect: bool = False, + starting_after: int = 0, + ) -> Iterator[str]: + min_seq = -1 if starting_after == 0 else starting_after + offset = 0 if starting_after < self.SETUP_N - 1 else starting_after - self.SETUP_N + 2 + emit_setup = not reconnect or starting_after < self.SETUP_N - 1 + yield from self._translate(request, meta, offset=offset, min_seq=min_seq, emit_setup=emit_setup) + + def _finalize_session(self) -> None: + self._del_meta() + if self.store.exists(): + self.store.cleanup() + + def _meta_key(self) -> str: + return f'sse:openai:meta:{self.store.user_uuid}:{self.store.idempotency_key}' + + def _save_meta(self, meta: dict) -> None: + self.store.redis_client.set(self._meta_key(), orjson.dumps(meta), ex=settings.SSE_STREAM_TTL) + + def _del_meta(self) -> None: + self.store.redis_client.delete(self._meta_key()) + + @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(chunk: str) -> tuple[str, dict, int | None]: + event, data, event_id = '', {}, None + for line in chunk.split('\n'): + if line.startswith('id:'): + event_id = 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, event_id + + @classmethod + def _token_seq(cls, event_id: int) -> int: + return cls.SETUP_N + (event_id - 2) + + def _response(self, meta: dict, status: str, output: list): + return { + 'id': meta['rid'], + 'object': 'response', + 'created_at': meta['ts'], + 'model': meta['model'], + 'status': status, + 'output': output, + 'store': True, + 'text': {'format': {'type': 'text'}}, + } + + def _setup(self, meta: dict, min_seq: int) -> Iterator[str]: + item = {'id': meta['mid'], 'type': 'message', 'status': 'in_progress', 'role': 'assistant', 'content': []} + created = self._response(meta, 'in_progress', []) + for seq, (event_type, fields) in enumerate(( + ('response.created', {'response': created}), + ('response.output_item.added', {'item': item, **self.IDX}), + ('response.content_part.added', { + 'item_id': meta['mid'], + 'part': {'type': 'output_text', 'text': '', 'annotations': []}, + **self.IDX, + }), + )): + if seq > min_seq: + yield self._sse(event_type, seq, **fields) + + def _closing(self, meta: dict, text: str, first_seq: int, min_seq: int) -> Iterator[str]: + part = {'type': 'output_text', 'text': text, 'annotations': []} + item = {'id': meta['mid'], 'type': 'message', 'status': 'completed', 'role': 'assistant', 'content': [part]} + ctx = {'item_id': meta['mid'], **self.IDX} + seq = first_seq + for event_type, fields in ( + ('response.output_text.done', {'text': text, **ctx}), + ('response.completed', {'response': self._response(meta, 'completed', [item])}), + ): + if seq > min_seq: + yield self._sse(event_type, seq, **fields) + seq += 1 + + def _translate( + self, + request: HttpRequest | None, + meta: dict, + *, + offset: int, + min_seq: int, + emit_setup: bool, + ) -> Iterator[str]: + if emit_setup: + yield from self._setup(meta, min_seq) + tokens, ctx, last_seq = [], {'item_id': meta['mid'], **self.IDX}, min_seq + for chunk in SSEChatStreamService(self.store).event_stream(request=request, offset=offset): + if chunk == SSEChatStreamService.HEARTBEAT: + continue + event, data, event_id = self._parse_chunk(chunk) + if event == 'token' and event_id and (token := data.get('content', '')): + tokens.append(token) + seq = self._token_seq(event_id) + if seq > min_seq: + yield self._sse('response.output_text.delta', seq, delta=token, **ctx) + last_seq = seq + elif event == 'done': + text = data.get('content') or ''.join(tokens) + close_seq = self._token_seq(event_id - 1) + 1 if event_id else self.SETUP_N + yield from self._closing(meta, text, close_seq, min_seq) + self._finalize_session() + return + elif event == 'error': + err_seq = last_seq + 1 + if err_seq > min_seq: + yield self._sse( + 'response.failed', + err_seq, + response={ + 'id': meta['rid'], + 'object': 'response', + 'status': 'failed', + 'error': {'code': 'server_error', 'message': str(data.get('detail', ''))}, + }, + ) + self._finalize_session() + return @@ -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 @@ -106,4 +111,7 @@ BUILDKIT_PROGRESS=plain # SERVER DJANGO_RUNSERVER_HIDE_WARNING=true -PYTHONWARNINGS=ignore::UserWarning:polymorphic # temporarily \ No newline at end of file +PYTHONWARNINGS=ignore::UserWarning:polymorphic # temporarily + +# SSE STREAMING +FF__STREAMING_ENABLED=True \ No newline at end of file @@ -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 ]