@@ -488,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) @@ -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]: @@ -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 @@ -1,6 +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 @@ -147,7 +148,7 @@ class NeuronModel(BaseModel, OrderedModel): @property def streaming(self) -> bool: - return hasattr(self.service, 'make_stream') + return settings.FF__STREAMING_ENABLED and hasattr(self.service, 'make_stream') def __str__(self): return self.title @@ -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) @@ -111,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