@@ -1,7 +1,10 @@ +from django.utils.translation import gettext as _ from ninja import Router from ninja.errors import HttpError +from rest_framework_simplejwt.exceptions import TokenError +from rest_framework_simplejwt.tokens import RefreshToken -from authentication.schemas import UserSchema +from authentication.schemas import UserSchema, RefreshInSchema, AccessOutSchema from authentication.security import SyncAuthBearer from authentication.selectors.user_selector import UserSelector @@ -15,3 +18,12 @@ def get_user_data(request): return UserSelector.detail(user=request.auth, provider=request.provider) except Exception as exc: raise HttpError(400, f'{exc}') + + +@router.post('refresh', auth=None, tags=['auth/refresh'], response=AccessOutSchema) +def refresh_token(request, payload: RefreshInSchema): + try: + refresh = RefreshToken(payload.refresh) + return AccessOutSchema(access=str(refresh.access_token)) + except TokenError: + raise HttpError(401, _('Invalid or expired refresh token')) @@ -43,3 +43,11 @@ class UserSchema(Schema): def resolve_account_type(obj: CustomUserModel): return obj.account_type + +class RefreshInSchema(Schema): + refresh: str + + +class AccessOutSchema(Schema): + access: str + @@ -67,7 +67,6 @@ urlpatterns = [ views.UpdateProfilePictureAPIView.as_view(), name='update-profile-pic', ), - path('token/refresh', TokenRefreshView.as_view(), name='refresh-jwt'), path( 'business-host', views.BusinessHostAPIView.as_view(), @@ -639,6 +639,14 @@ msgid "The model is currently disabled. Please try again later." msgstr "" "Модель в настоящее время неактивна. Пожалуйста, повторите попытку позже." +#: ml_model/exceptions.py:119 +msgid "Service is currently unavailable due to high demand. Please try again later" +msgstr "Сервис временно недоступен из-за высокой нагрузки. Пожалуйста, попробуйте позже" + +#: ml_model/exceptions.py:135 +msgid "Image analysis error. Please try another image." +msgstr "Ошибка анализа изображения. Попробуйте другую картинку." + #: ml_model/exceptions.py:23 msgid "Your request was blocked by our moderation system" msgstr "Ваш запрос был заблокирован нашей системой модерации" @@ -672,6 +680,10 @@ msgstr "" "Формат вложенного файла не поддерживается. Доступные форматы: " "%(available_extensions)s." +#: ml_model/exceptions.py:60 +msgid "The file may be corrupted. Please try another one." +msgstr "Возможно, файл повреждён. Попробуйте загрузить другой файл." + #: ml_model/exceptions.py:58 msgid "The length of the context has been exceeded." msgstr "Длина контекста превышена." @@ -1020,6 +1032,9 @@ msgstr "" "Вы можете осуществлять поиск по e-mail пользователя, точному названию " "компании" +msgid "Image is ready" +msgstr "Изображение готово" + #: payments/admin.py:40 payments/admin.py:100 msgid "Missing" msgstr "Отсутствующий" @@ -1283,6 +1298,10 @@ msgstr "" "Случилась ошибка во время генерации. Она может возникать из-за того, что " "NSFW-контент запрещен. Попробуйте снова" +#: tools/media/apis.py:168 +msgid "Temporary issues with the service, we are already working on a solution." +msgstr "Временные неполадки с сервисом, мы уже работаем над их решением." + #: tools/chats/apis.py:244 msgid "The message has already been deleted" msgstr "Сообщение уже было удалено" @@ -1354,6 +1373,10 @@ msgstr "Отсутствует обязательный параметр: 'messa msgid "Model not found" msgstr "Модель не найдена" +#: authentication/routes/v1.py:28 +msgid "Invalid or expired refresh token" +msgstr "Неверный или истёкший refresh токен" + #~ msgid "Regular users cannot send introductory letters" #~ msgstr "Обычные пользователи не могут отсылать письма" @@ -33,8 +33,9 @@ class EmbeddingService: redis_client: redis.Redis, message_uid: str, chunk_id: int, + model: str = 'text-embedding-3-large', ) -> int: - embedding, e_total_tokens = cls._get_embedding(client=client, content=chunk) + embedding, e_total_tokens = cls._get_embedding(client=client, content=chunk, model=model) cls._save_embeddings( redis_client=redis_client, message_uid=message_uid, @@ -45,8 +46,13 @@ class EmbeddingService: return e_total_tokens @classmethod - def _get_embedding(cls, client: httpx.Client, content: str) -> Tuple[List[float], int]: - response = client.post(url='embeddings', json={'model': 'text-embedding-3-large', 'input': content}) + def _get_embedding( + cls, + client: httpx.Client, + content: str, + model: str = 'text-embedding-3-large', + ) -> Tuple[List[float], int]: + response = client.post(url='embeddings', json={'model': model, 'input': content}) response.raise_for_status() data = response.json() return data['data'][0]['embedding'], data['usage']['total_tokens'] @@ -72,6 +78,7 @@ class EmbeddingService: message_uid: str, user_query_embeddings: List[float], top_k: int = 10, + index_name: str = 'ml_model-index', ) -> List[Document]: base_query = ( f'@message_uid:{{{message_uid}}}=>[KNN {top_k} @section_embeddings $vector AS vector_score]' @@ -84,7 +91,7 @@ class EmbeddingService: .dialect(2) ) params_dict = {'vector': np.array(user_query_embeddings).astype(dtype=np.float32).tobytes()} - results = redis_client.ft('ml_model-index').search(query, params_dict) + results = redis_client.ft(index_name).search(query, params_dict) return results.docs @classmethod @@ -97,7 +104,16 @@ class EmbeddingService: """ @classmethod - def get_large_file_data(cls, msg_uid: UUID4, chunks, proxy, user_content): + def get_large_file_data( + cls, + msg_uid: UUID4, + chunks, + proxy, + user_content, + model: str = 'text-embedding-3-large', + index_name: str = 'ml_model-index', + top_k: int = 10, + ): redis_client = redis.Redis(host=settings.REDIS_HOST, port=settings.REDIS_PORT, db=0) embedding_tokens = 0 message_uid = str(msg_uid).replace('-', '_') @@ -112,17 +128,29 @@ class EmbeddingService: for chunk_id, chunk in enumerate(chunks): threads.append( executor.submit( - cls.process_chunk, client, chunk, redis_client, message_uid, chunk_id + cls.process_chunk, + client, + chunk, + redis_client, + message_uid, + chunk_id, + model, ) ) for thread in as_completed(threads): embedding_tokens += thread.result() - query_embedding, e_total_tokens = cls._get_embedding(client=client, content=user_content) + query_embedding, e_total_tokens = cls._get_embedding( + client=client, content=user_content, model=model + ) embedding_tokens += e_total_tokens result = [ s['section_text'] for s in cls.search_via_embeddings( - redis_client=redis_client, message_uid=message_uid, user_query_embeddings=query_embedding + redis_client=redis_client, + message_uid=message_uid, + user_query_embeddings=query_embedding, + top_k=top_k, + index_name=index_name, ) ] drop_redis_vectors.delay(message_uid) @@ -1,5 +1,6 @@ from ml_model.services.chatgpt import Chatgpt from ml_model.services.chatgpt_5 import Chatgpt_5 +from ml_model.services.chatgpt_5_4 import Chatgpt_5_4 from ml_model.services.claude import Claude from ml_model.services.codellama import Codellama from ml_model.services.dalle import Dalle @@ -7,18 +8,21 @@ from ml_model.services.deepl import Deepl from ml_model.services.deepseek import Deepseek from ml_model.services.djourney import Djourney from ml_model.services.epicphotogasm import Epicphotogasm +from ml_model.services.elevenlabs import Elevenlabs from ml_model.services.flux import Flux from ml_model.services.flux_2 import Flux_2 from ml_model.services.fluxkrea import Fluxkrea from ml_model.services.fluxlorafast import Fluxlorafast from ml_model.services.fluxproultra import Fluxproultra from ml_model.services.gemini import Gemini +from ml_model.services.gemini_3_1 import Gemini_3_1 from ml_model.services.gemma import Gemma from ml_model.services.geminiimage import Geminiimage from ml_model.services.gptimage import Gptimage from ml_model.services.granite import Granite from ml_model.services.grok import Grok from ml_model.services.grok_4_1_fast import Grok_4_1_Fast +from ml_model.services.grok_imagine_video import Grok_Imagine_Video from ml_model.services.hailuo import Hailuo from ml_model.services.hunyuan import Hunyuan from ml_model.services.iconic import Iconic @@ -31,6 +35,7 @@ from ml_model.services.lightning import Lightning from ml_model.services.llama import Llama from ml_model.services.logoai import Logoai from ml_model.services.lyria import Lyria +from ml_model.services.ltx import Ltx from ml_model.services.midjourney import Midjourney from ml_model.services.minimaxvideo import Minimaxvideo from ml_model.services.minimaxmusic import Minimaxmusic @@ -38,11 +43,15 @@ from ml_model.services.minimaxmusic_lite import Minimaxmusic_Lite from ml_model.services.mistral import Mistral from ml_model.services.musicgen import Musicgen from ml_model.services.nanobanana import Nanobanana +from ml_model.services.nanobanana_2 import Nanobanana_2 from ml_model.services.perplexity import Perplexity from ml_model.services.pulid import Pulid from ml_model.services.photon import Photon +from ml_model.services.prunaai import Prunaai +from ml_model.services.pruna_v import Pruna_V from ml_model.services.qwen import Qwen from ml_model.services.qwen_235B import Qwen_235B +from ml_model.services.qwen_3_5 import Qwen_3_5 from ml_model.services.qwen_3_max_thinking import Qwen_3_Max_Thinking from ml_model.services.raifgpt import Raifgpt from ml_model.services.ray import Ray @@ -1,17 +1,8 @@ import base64 import itertools import logging -import re -import subprocess import time -import zipfile -from concurrent.futures import ThreadPoolExecutor, as_completed - -import numpy as np -import openpyxl -import fitz -import redis from django.utils.translation import gettext_lazy as _ from datetime import timedelta @@ -20,7 +11,6 @@ from io import BufferedReader, BytesIO from math import ceil from typing import Generator, List, Optional, Dict, Any, Tuple -import docx2txt import filetype import httpx import tiktoken @@ -35,21 +25,16 @@ from langchain_core.messages import ( from langchain_core.prompts.prompt import PromptTemplate from langchain_core.runnables import RunnableWithMessageHistory from langchain_openai.chat_models import ChatOpenAI -from langchain_text_splitters import RecursiveCharacterTextSplitter from PIL import Image -from redis.commands.search.document import Document -from redis.commands.search.query import Query from backend import settings from messages.models import BaseStore, Message from ml_model.constants import TEMPORARY_TEST_TEXT -from ml_model.exceptions import FileExtensionNotSupported -from ml_model.models import ( - ModelConfiguration, - NeuronModel -) +from ml_model.exceptions import FileExtensionNotSupported, CorruptedFileError +from ml_model.models import ModelConfiguration, NeuronModel +from ml_model.services.EmbeddingService import EmbeddingService +from ml_model.services.FileService import FileProcessingService from ml_model.services.base import SimpleService -from ml_model.tasks import drop_redis_vectors from payments.exceptions.insufficient_balance import InsufficientBalance from payments.selectors.payment_plan_selector import PaymentPlanSelector from poller.models import Proxy @@ -73,32 +58,30 @@ class Chatgpt(SimpleService): 'input': Decimal('0.0003'), 'output': Decimal('0.0003'), 'web_search': { - 'low': Decimal('12.5'), # 1 call - 'medium': Decimal('13.75'), # 1 call - 'high': Decimal('15') # 1 call - } + 'low': Decimal('12.5'), # 1 call + 'medium': Decimal('13.75'), # 1 call + 'high': Decimal('15'), # 1 call + }, }, 'gpt-4o': { 'input': Decimal('0.005'), 'output': Decimal('0.005'), 'web_search': { - 'low': Decimal('15'), # 1 call - 'medium': Decimal('17.5'), # 1 call - 'high': Decimal('25') # 1 call - } + 'low': Decimal('15'), # 1 call + 'medium': Decimal('17.5'), # 1 call + 'high': Decimal('25'), # 1 call + }, }, - 'gpt-oss-120b': { - 'input': Decimal('0.0002'), - 'output': Decimal('0.0002') - } + 'gpt-oss-120b': {'input': Decimal('0.0002'), 'output': Decimal('0.0002')}, } TOOLS_TOKEN_COSTS = { - 'text-embedding-3-large': { - 'output': Decimal('0.000065') - } + 'text-embedding-3-small': {'output': Decimal('0.00001')}, + 'text-embedding-3-large': {'output': Decimal('0.000065')}, } + EMBEDDING_MODEL_FOR_BILLING = 'text-embedding-3-small' + TOKEN_LIMITS = { 'o3-mini': 100_000, 'gpt-4o-mini': 64_000, @@ -113,7 +96,7 @@ class Chatgpt(SimpleService): @property def neuron_model(self): - return NeuronModel.objects.get(title='ChatGPT') + return NeuronModel.objects.get(title='ChatGPT 4') def make( self, @@ -131,19 +114,24 @@ class Chatgpt(SimpleService): normalized_image = None embedding_tokens = 0 chunks = [] + text_chunks: list[str] = [] if file: - try: - file_bytes = input_message.file.read() - kind = filetype.guess(file_bytes[:20]) - raw_file_extension = kind.extension - file_extension = self._get_file_extension(raw_file_extension, file_bytes) - if file_extension in ('pdf', 'doc', 'docx', 'xlsx'): - chunks = self._get_file_data(file_extension, file_bytes) - else: - image = file - normalized_image, image_size, image_data = self._get_image_data(file_bytes, file_extension) - input_content.append(image_data) - except: + file_service = FileProcessingService + file_bytes = input_message.file.read() + kind = filetype.guess(file_bytes[:20]) + if not kind: + 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) + text_chunks = EmbeddingService.split_text_to_chunks(text) + chunks = [HumanMessage(content=chunk_text) for chunk_text in text_chunks] + elif file_extension in ('jpg', 'jpeg', 'png', 'webp'): + image = file + normalized_image, image_size, image_data = self._get_image_data(file_bytes, file_extension) + input_content.append(image_data) + else: raise FileExtensionNotSupported(['PDF', 'DOC', 'DOCX', 'XLSX', 'JPG', 'JPEG', 'PNG', 'WEBP']) for proxy in Proxy.objects.all(): self.llm = ChatOpenAI( @@ -165,7 +153,9 @@ class Chatgpt(SimpleService): get_session_history=lambda _: chat_history, ) llm_input = [SystemMessage(content=user_system_prompt), HumanMessage(content=input_content)] - input_tokens, input_embedding_tokens = self._get_input_tokens(file, image, chunks, chat_history, llm_input) + input_tokens, input_embedding_tokens = self._get_input_tokens( + file, image, chunks, chat_history, llm_input, model_name + ) output_tokens = 0 self.assert_enough_balance( input_tokens, image_size, model=self.llm.model_name, embedding_tokens=input_embedding_tokens @@ -173,28 +163,44 @@ class Chatgpt(SimpleService): if model_name == 'gpt-oss-120b': system = chat_history.messages.pop(0) messages = [ - {'role': 'user' if isinstance(msg, HumanMessage) else 'assistant', 'content': msg.content} + { + 'role': 'user' if isinstance(msg, HumanMessage) else 'assistant', + 'content': msg.content, + } for msg in chat_history.messages ] messages.insert(0, {'role': 'system', 'content': system.content}) messages.insert(0, {'role': 'system', 'content': user_system_prompt}) 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 = self.get_large_file_data(chunks, proxy, input_message.content) - messages[-1]['content'] = self.make_embeddings_prompt( - document_name=document_name, section_texts=file_data, question=input_message.content + 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'Используй системный промпт. Содержание файла: ' - f'{chunks}. Вопрос: {input_message.content}') - json_data = { - 'model': f'openai/{model_name}', - 'messages': messages - } + messages[-1]['content'] = ( + 'Используй системный промпт. Содержание файла: ' + f'{"".join(text_chunks)}. Вопрос: {input_message.content}' + ) + json_data = {'model': f'openai/{model_name}', 'messages': messages} response = httpx.post( - url='https://openrouter.ai/api/v1/chat/completions', proxy=f'{proxy.protocol}://{proxy.address}', - headers={'Authorization': f'Bearer {settings.OPENROUTER_API_KEY}'}, timeout=600, json=json_data + url='https://openrouter.ai/api/v1/chat/completions', + proxy=f'{proxy.protocol}://{proxy.address}', + headers={'Authorization': f'Bearer {settings.OPENROUTER_API_KEY}'}, + timeout=600, + json=json_data, ) if ( (data := response.json()) @@ -212,7 +218,10 @@ class Chatgpt(SimpleService): elif model_name == 'o3-mini': system = chat_history.messages.pop(0) messages = [ - {'role': 'user' if isinstance(msg, HumanMessage) else 'assistant', 'content': msg.content} + { + 'role': 'user' if isinstance(msg, HumanMessage) else 'assistant', + 'content': msg.content, + } for msg in chat_history.messages ] messages.insert(0, {'role': 'system', 'content': system.content}) @@ -225,13 +234,24 @@ class Chatgpt(SimpleService): elif file: 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 = self.get_large_file_data(chunks, proxy, input_message.content) - messages[-1]['content'] = self.make_embeddings_prompt( - document_name=document_name, section_texts=file_data, question=input_message.content + 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'Используй системный промпт. Содержание файла: ' - f'{chunks}. Вопрос: {input_message.content}') + messages[-1]['content'] = ( + 'Используй системный промпт. Содержание файла: ' + f'{"".join(text_chunks)}. Вопрос: {input_message.content}' + ) json_data = { 'model': model_name, 'messages': messages @@ -246,27 +266,38 @@ class Chatgpt(SimpleService): messages.insert(0, {'role': 'system', 'content': system.content}) messages.insert(0, {'role': 'system', 'content': user_system_prompt}) search_context_size, json_data = self.get_web_search_data( - info.get('web_search', 'Средний контекст'), - model_name, - messages + info.get('web_search', 'Средний контекст'), model_name, messages ) info['web_search'] = search_context_size if image: messages[-1]['content'] = [ {'type': 'input_text', 'text': input_message.content}, - {'type': 'input_image', 'image_url': image_data['image_url']['url']} + {'type': 'input_image', 'image_url': image_data['image_url']['url']}, ] elif file: 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 = self.get_large_file_data(chunks, proxy, input_message.content) - messages[-1]['content'] = self.make_embeddings_prompt( - document_name=document_name, section_texts=file_data, question=input_message.content + 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'Используй системный промпт. Содержание файла: ' - f'{chunks}. Вопрос: {input_message.content}') - input_tokens, output_tokens, response = self.call_openai_api(proxy=proxy, endpoint='responses',json_data=json_data) + messages[-1]['content'] = ( + 'Используй системный промпт. Содержание файла: ' + f'{"".join(text_chunks)}. Вопрос: {input_message.content}' + ) + input_tokens, output_tokens, response = self.call_openai_api( + proxy=proxy, endpoint='responses', json_data=json_data + ) elif image: response = self.llm.invoke(llm_input) chat_history.add_ai_message(response) @@ -274,12 +305,23 @@ class Chatgpt(SimpleService): input_tokens = self.count_text_tokens([*chat_history.messages]) 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 = self.get_large_file_data(chunks, proxy, input_message.content) + 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', + ) user_input = [ SystemMessage(content=user_system_prompt), - HumanMessage(self.make_embeddings_prompt( - document_name=document_name, section_texts=file_data, question=input_message.content - )) + HumanMessage( + EmbeddingService.make_embeddings_prompt( + document_name=document_name, + section_texts=file_data, + question=input_message.content, + ) + ), ] input_tokens += self.count_text_tokens(user_input) response = conversation.invoke( @@ -290,9 +332,11 @@ class Chatgpt(SimpleService): input = [ SystemMessage(content=user_system_prompt), HumanMessage( - content=f'Используй системный промпт. Содержание файла: ' - f'{chunks}. Вопрос: {input_message.content}' - ) + content=( + 'Используй системный промпт. Содержание файла: ' + f'{"".join(text_chunks)}. Вопрос: {input_message.content}' + ) + ), ] input_tokens += self.count_text_tokens(input) response = conversation.invoke( @@ -385,7 +429,8 @@ class Chatgpt(SimpleService): input_tokens: int, image_size: tuple | None, model: str = 'gpt-3.5-turbo', - embedding_tokens: int = 0 + embedding_tokens: int = 0, + output_tokens: int = 0, ): balance = PaymentPlanSelector(self.store.user).get_current_balance() total_tokens = input_tokens @@ -393,9 +438,13 @@ class Chatgpt(SimpleService): total_tokens += self.count_image_tokens(image_size) input_cost = self.TOKENS_COST[model]['input'] * total_tokens if embedding_tokens > 0: - input_cost += embedding_tokens * self.TOOLS_TOKEN_COSTS['text-embedding-3-large']['output'] - if input_cost > balance: - raise InsufficientBalance(balance, input_cost) + input_cost += ( + embedding_tokens + * self.TOOLS_TOKEN_COSTS[self.EMBEDDING_MODEL_FOR_BILLING]['output'] + ) + output_cost = self.TOKENS_COST[model]['output'] * output_tokens + if input_cost + output_cost > balance: + raise InsufficientBalance(balance, input_cost + output_cost) def calculate_price( self, @@ -416,7 +465,10 @@ class Chatgpt(SimpleService): if info.get('code_interpreter', False): price += self.TOKENS_COST[model]['code_interpreter'] if embedding_tokens > 0: - price += self.TOOLS_TOKEN_COSTS['text-embedding-3-large']['output'] * embedding_tokens + price += ( + self.TOOLS_TOKEN_COSTS[self.EMBEDDING_MODEL_FOR_BILLING]['output'] + * embedding_tokens + ) return price.quantize(Decimal('0.1'), rounding='ROUND_UP') def count_image_tokens(self, image_size: tuple, model_version: str = 'gpt-4o') -> int: @@ -461,20 +513,6 @@ class Chatgpt(SimpleService): return total_tokens - def _get_file_extension(self, raw_file_extension: str, file_bytes: bytes) -> str: - if raw_file_extension == 'zip': - signatures = { - 'xlsx': 'xl/workbook.xml', - 'docx': 'word/document.xml' - } - with zipfile.ZipFile(BytesIO(file_bytes), 'r') as zip_file: - namelist = zip_file.namelist() - for format_name, required_file in signatures.items(): - if required_file in namelist: - return format_name - raise - return raw_file_extension - def _get_image_data(self, file_bytes: bytes, file_extension: str) -> Tuple: normalized_image = Image.open(BytesIO(file_bytes)).convert('RGB') buf = BytesIO() @@ -486,16 +524,7 @@ class Chatgpt(SimpleService): image_data = {'type': 'image_url', 'image_url': {'url': image_url}} return normalized_image, image_size, image_data - def _get_file_data(self, file_extension: str, file_bytes: bytes) -> list[HumanMessage]: - is_word = file_extension in ('doc', 'docx') - method_name = 'word' if is_word else file_extension - operation = getattr(self, f'get_{method_name}_data') - text = operation(file_extension, file_bytes) if is_word else operation(file_bytes) - if file_extension != 'xlsx': - text = re.sub(r'\n{2,}', '\n', text) - return self.split_text_to_chunks(text) - - def _get_input_tokens(self, file, image, chunks, chat_history, llm_input): + def _get_input_tokens(self, file, image, chunks, chat_history, llm_input, model_name=None): input_embedding_tokens = 0 if file and not image: if sum([len(chunk.content) for chunk in chunks]) > 20_000: @@ -503,49 +532,17 @@ class Chatgpt(SimpleService): input_embedding_tokens = len(chunks) * 600 else: input_tokens = self.count_text_tokens([*chat_history.messages, *llm_input, *chunks]) - elif image: + elif image and model_name in ('gpt-4o', 'gpt-4o-mini'): input_tokens = self.count_text_tokens(llm_input) else: input_tokens = self.count_text_tokens([*chat_history.messages, *llm_input]) return input_tokens, input_embedding_tokens - def get_large_file_data(self, chunks, proxy, user_content): - redis_client = redis.Redis(host=settings.REDIS_HOST, port=settings.REDIS_PORT, db=0) - embedding_tokens = 0 - message_uid = str(self.store.messages.first().pk).replace('-', '_') - 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=600, - ) as client: - threads = [] - with ThreadPoolExecutor(max_workers=settings.MAX_THREADS) as executor: - for chunk_id, chunk in enumerate(chunks): - threads.append( - executor.submit(self.process_chunk, client, chunk, redis_client, message_uid, chunk_id) - ) - for thread in as_completed(threads): - embedding_tokens += thread.result() - query_embedding, e_total_tokens = self.get_embedding(client=client, content=user_content) - embedding_tokens += e_total_tokens - result = [ - s['section_text'] - for s in self.search_via_embeddings( - redis_client=redis_client, - message_uid=message_uid, - user_query_embeddings=query_embedding - ) - ] - drop_redis_vectors.delay(message_uid) - redis_client.close() - return embedding_tokens, result - def get_web_search_data(self, search_size: str, model_name: str, messages: List[Dict[str, any]]): search_context_sizes = { 'Малый контекст': 'low', 'Средний контекст': 'medium', - 'Большой контекст': 'high' + 'Большой контекст': 'high', } search_context_size = search_context_sizes.get(search_size) json_data = { @@ -555,202 +552,23 @@ class Chatgpt(SimpleService): { 'type': 'web_search_preview', 'search_context_size': search_context_size, - 'user_location': {'type': 'approximate', 'country': 'RU'} + 'user_location': {'type': 'approximate', 'country': 'RU'}, } - ] + ], } return search_context_size, json_data - def get_pdf_data(self, pdf_data: bytes) -> str: - """ - Extracting text from pdf-file - :param pdf_file: uploaded pdf file - :return: pdf-file content - """ - try: - doc = fitz.open(stream=pdf_data, filetype="pdf") - raw_text = '' - for page_number, page in enumerate(doc, start=1): - content = page.get_text("text") - if content: - raw_text += content - doc.close() - fitz.TOOLS.store_shrink(100) - except Exception: - return f"Ошибка: Файл поврежден или не может быть прочитан." - return f'Содержимое файла: {raw_text.strip()}' - - def get_xlsx_data(self, xlsx_data: bytes) -> str: - """ - Extracting text from xlsx-file - :param xlsx_file: uploaded xlsx file - :return: xlsx_file content - """ - try: - xlsx_content = BytesIO(xlsx_data) - workbook = openpyxl.load_workbook(xlsx_content) - raw_text = '' - for sheet_name in workbook.sheetnames: - sheet = workbook[sheet_name] - for row in sheet.iter_rows(values_only=True): - raw_text += f'Данные ряда: {row}\n' - except Exception: - raw_text = 'Произошла ошибка во время чтения файла' - return f'Содержимое файла: {raw_text}' - - def get_word_data(self, extension: str, word_data: bytes) -> str: - """ - Extracting text from word-file - :param extension: extension of uploaded word file - :param word_file: uploaded word file - :return: word-file content - """ - try: - if extension == 'docx': - text = docx2txt.process(BytesIO(word_data)) - elif extension == 'doc': - process = subprocess.Popen( - ['antiword', '-w', '0', '-'], - stdin=subprocess.PIPE, - stdout=subprocess.PIPE, - stderr=subprocess.PIPE, - ) - text, _ = process.communicate(input=word_data) - text = text.decode('utf-8') - else: - text = '' - except Exception: - text = 'Файл поврежден или не может быть прочитан.' - if text.strip(): - return f'Это текст, извлечённый из загруженного WORD-файла:\n{text}' - else: - return 'Файл пуст или содержит изображения, из которых невозможно извлечь текст.' - - def split_text_to_chunks( - self, raw_text: str, chunk_size: int = 4000, overlap: int = 200 - ) -> list[HumanMessage]: - """ - Splitting file raw text to chunks - :param raw_text: full text which file includes - :param chunk_size: еhe maximum size of each chunk - :param overlap: еhe number of overlapping characters between chunks - :return: list of chunks - """ - text_splitter = RecursiveCharacterTextSplitter( - chunk_size=chunk_size, chunk_overlap=overlap, length_function=len, separators=["\n\n", "\n", ".", " ", ""] - ) - chunks = text_splitter.split_text(raw_text) - return [HumanMessage(chunk) for chunk in chunks] - - def process_chunk( - self, client: httpx.Client, chunk: HumanMessage, redis_client: redis.Redis, message_uid:str, chunk_id: int - ) -> int: - ''' - A method for getting and saving embeddings from a single chunk - :param client: Httpx client - :param chunk: a HumanMessage object with a content as a part of a full text - :param redis_client: Redis client - :param message_uid: UID of user's message - :param chunk_id: a sequence number of a chunk - ''' - embedding, e_total_tokens = self.get_embedding(client=client, content=chunk.content) - self.save_embeddings( - redis_client=redis_client, - message_uid=message_uid, - chunk_id=chunk_id, - text=chunk.content, - embeddings=embedding - ) - return e_total_tokens - - def get_embedding(self, client: httpx.Client, content: str) -> Tuple[List[float], int]: - ''' - A method for converting raw text (content) into embeddings - using OpenAI API request - :param client: Httpx client - :param content: raw text of a chunk - ''' - response = client.post( - url="embeddings", - json={ - 'model': 'text-embedding-3-large', - 'input': content - } - ) - response.raise_for_status() - data = response.json() - return data['data'][0]['embedding'], data['usage']['total_tokens'] - - def save_embeddings( - self, redis_client: redis.Redis, message_uid: str, chunk_id: int, text: str, embeddings: List[float] - ) -> None: - ''' - A method for saving embeddings in Redis - :param redis_client: Redis client - :param message_uid: UID of user's message - :param chunk_id: a sequence number of a chunk - :param text: a chunk content - :param embeddings: a list of embeddings getting from a chunk - ''' - embeddings_bytes = np.array(embeddings).astype(dtype=np.float32).tobytes() - redis_client.hset( - f'ml_model:messages:{message_uid}:vectors:{chunk_id}', - mapping={ - 'message_uid': message_uid, - 'section_text': text, - 'section_embeddings': embeddings_bytes - } - ) - - def search_via_embeddings( - self, redis_client: redis.Redis, message_uid: str, user_query_embeddings: List[float], top_k: int = 10 - ) -> List[Document]: - ''' - A method for searching similar vectors to user's query - :param redis_client: Redis client - :param message_uid: UID of user's message - :param user_query_embeddings: a list of embeddings getting from user's query - :param top_k: a number of max return documents - ''' - base_query = f'@message_uid:{{{message_uid}}}=>[KNN {top_k} @section_embeddings $vector AS vector_score]' - query = ( - Query(base_query) - .return_fields('section_text') - .sort_by("vector_score") - .paging(0, top_k) - .dialect(2) - ) - params_dict = {"vector": np.array(user_query_embeddings).astype(dtype=np.float32).tobytes()} - results = redis_client.ft('ml_model-index').search(query, params_dict) - return results.docs - - def make_embeddings_prompt(self, document_name: str, section_texts: List[str], question: str) -> str: - ''' - A method for making a prompt using found embeddings - :param document_name: name of the loaded document - :param section_texts: list of sections' contents - :param question: user question - ''' - return f"""Ты — аналитик данных. Отвечай только на основе предоставленного контекста. - Название файла: {document_name} - Фрагменты: - { - '\n'.join(section_texts) - } - Вопрос: {question} - """ - def call_openai_api( - self, proxy: Proxy, endpoint: str, json_data: Dict[str, Any] + self, proxy: Proxy, endpoint: str, json_data: Dict[str, Any] ) -> Tuple[Any, Any, AIMessage] | Tuple[List[float], int]: - ''' + """ A method for sending a request to official openai API :param proxy: Proxy settings object with protocol and address. :param endpoint: Str URL part for the OpenAI API request :param json_data: Payload for the OpenAI API request :return: Tuple of (input_tokens, output_tokens, AIMessage instance with response content) :raises: Exception: If the response is invalid or incomplete - ''' + """ with httpx.Client( base_url='https://api.openai.com/v1', proxy=f'{proxy.protocol}://{proxy.address}', @@ -775,17 +593,30 @@ class Chatgpt(SimpleService): output_tokens = resp.json()['usage']['completion_tokens'] response = AIMessage(content=content) return input_tokens, output_tokens, response + elif ( + endpoint == 'responses' + and (data := resp.json()) + and data.get('output') + and (content := data['output'][-1]['content'][0]['text']) + ): + input_tokens = resp.json()['usage']['input_tokens'] + output_tokens = resp.json()['usage']['output_tokens'] + response = AIMessage(content=content.replace('\\n', '\n')) + return input_tokens, output_tokens, response elif ( endpoint == 'responses' and (data := resp.json()) and data.get('output') and ( - content := data['output'][-1]['content'][0]['text'] + image := next( + (item['result'] for item in data['output'] if item.get('result')), + None, + ) ) ): input_tokens = resp.json()['usage']['input_tokens'] output_tokens = resp.json()['usage']['output_tokens'] - response = AIMessage(content=content.replace('\\n', '\n')) + response = AIMessage(content=[{'generate_image': True, 'image': image}]) return input_tokens, output_tokens, response else: raise Exception('GPT not answer correctly, please retry later') @@ -6,9 +6,11 @@ import filetype from langchain_core.messages import HumanMessage, SystemMessage from messages.models import Message -from ml_model.exceptions import FileExtensionNotSupported +from ml_model.exceptions import FileExtensionNotSupported, CorruptedFileError from ml_model.models import NeuronModel from ml_model.services import Chatgpt +from ml_model.services.EmbeddingService import EmbeddingService +from ml_model.services.FileService import FileProcessingService from poller.models import Proxy @@ -103,24 +105,29 @@ class Chatgpt_5(Chatgpt): image_size = None embedding_tokens = 0 chunks = [] + text_chunks: list[str] = [] if file: - try: - file_bytes = input_message.file.read() - kind = filetype.guess(file_bytes[:20]) - raw_file_extension = kind.extension - file_extension = self._get_file_extension(raw_file_extension, file_bytes) - if file_extension in ('pdf', 'doc', 'docx', 'xlsx'): - chunks = self._get_file_data(file_extension, file_bytes) - else: - image = file - _, image_size, image_data = self._get_image_data(file_bytes, file_extension) - except Exception: + file_service = FileProcessingService + file_bytes = input_message.file.read() + kind = filetype.guess(file_bytes[:20]) + if not kind: + 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) + text_chunks = EmbeddingService.split_text_to_chunks(text) + chunks = [HumanMessage(content=chunk_text) for chunk_text in text_chunks] + elif file_extension in ('jpg', 'jpeg', 'png', 'webp'): + image = file + _, image_size, image_data = self._get_image_data(file_bytes, file_extension) + else: raise FileExtensionNotSupported(['PDF', 'DOC', 'DOCX', 'XLSX', 'JPG', 'JPEG', 'PNG', 'WEBP']) chat_history = self.get_chat_history(model_name=model_name) chat_history.add_message(HumanMessage(content=input_message.content)) llm_input = [SystemMessage(content=user_system_prompt), HumanMessage(content=input_content)] input_tokens, input_embedding_tokens = self._get_input_tokens( - file, image, chunks, chat_history, llm_input + file, image, chunks, chat_history, llm_input, model_name ) self.assert_enough_balance( input_tokens, image_size, model=model_name, embedding_tokens=input_embedding_tokens @@ -143,16 +150,23 @@ class Chatgpt_5(Chatgpt): document_name = ( chunks[0].content.partition(f':{chr(10)}')[2].split(f'{chr(10)}')[0][:100] ) - embedding_tokens, file_data = self.get_large_file_data( - chunks, proxy, input_message.content + 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'] = self.make_embeddings_prompt( - document_name=document_name, section_texts=file_data, question=input_message.content + 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}' + 'Используй системный промпт. Содержание файла: ' + f'{"".join(text_chunks)}. Вопрос: {input_message.content}' ) json_data = { 'model': model_name, @@ -0,0 +1,284 @@ +import base64 +import time +from datetime import timedelta +from decimal import Decimal +from io import BytesIO + +import filetype +from django.core.files import File +from django.utils.translation import gettext_lazy +from langchain_core.messages import HumanMessage, SystemMessage, BaseMessage + +from messages.models import Message +from ml_model.exceptions import FileExtensionNotSupported, CorruptedFileError +from ml_model.models import NeuronModel +from ml_model.services import Chatgpt +from ml_model.services.EmbeddingService import EmbeddingService +from ml_model.services.FileService import FileProcessingService +from poller.models import Proxy + + +class Chatgpt_5_4(Chatgpt): + TOKENS_COST = { + 'gpt-5.4': { + 'input': Decimal('0.00125'), + 'output': Decimal('0.0075'), + 'web_search': { + 'low': Decimal('5'), # 1 call + 'medium': Decimal('5'), # 1 call + 'high': Decimal('5'), # 1 call + }, + 'code_interpreter': Decimal('15'), # 1 call + 'generated_image': Decimal('10.2'), + }, + 'gpt-5.4-pro': { + 'input': Decimal('0.015'), + 'output': Decimal('0.09'), + 'web_search': { + 'low': Decimal('5'), # 1 call + 'medium': Decimal('5'), # 1 call + 'high': Decimal('5'), # 1 call + }, + 'generated_image': Decimal('10.2'), + }, + } + + TOKEN_LIMITS = { + 'gpt-5.4': 1_050_000 // 2, + 'gpt-5.4-pro': 1_050_000 // 2, + } + + @property + def neuron_model(self): + return NeuronModel.objects.get(slug='chatgpt_5_4') + + def save_results( + self, + results: list[BaseMessage], + elapsed_time: timedelta, + generated_image: bytes | None, + save: bool = True, + ) -> list[Message]: + messages = [ + Message( + content=result.content, + elapsed_time=elapsed_time, + content_object=self.store, + file=File(BytesIO(generated_image), '.png') if generated_image else None, + ) + for result in results + ] + if save: + return Message.objects.bulk_create(messages) + return messages + + def calculate_price( + self, + input_tokens: int, + output_tokens: int, + model: str, + info: dict, + embedding_tokens: int = 0, + image: bool = False, + *args, + **kwargs, + ) -> Decimal: + price = ( + input_tokens * self.TOKENS_COST[model]['input'] + + output_tokens * self.TOKENS_COST[model]['output'] + ) + if info.get('web_search', 'Отключено') != 'Отключено': + price += self.TOKENS_COST[model]['web_search'].get(info.get('web_search', 'medium')) + if info.get('code_interpreter', False): + price += self.TOKENS_COST[model]['code_interpreter'] + if embedding_tokens > 0: + price += self.TOOLS_TOKEN_COSTS[self.EMBEDDING_MODEL_FOR_BILLING]['output'] * embedding_tokens + if image: + price += self.TOKENS_COST[model]['generated_image'] + 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 = info.pop('version', 'gpt-5.4') + user_system_prompt = info.pop('system_prompt', '') + plan_info = self.store.user.payment_plan + is_free_plan = plan_info and plan_info.plan.price <= 0 + if is_free_plan: + info.pop('web_search', None) + info.pop('code_interpreter', None) + info.pop('verbosity', None) + input_content = [{'type': 'text', 'text': input_message.content or ''}] + file = input_message.file + image = None + image_size = None + embedding_tokens = 0 + chunks = [] + text_chunks = [] + if file: + file_service = FileProcessingService + file_bytes = input_message.file.read() + kind = filetype.guess(file_bytes[:20]) + if not kind: + 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) + text_chunks = EmbeddingService.split_text_to_chunks(text) + chunks = [HumanMessage(content=chunk_text) for chunk_text in text_chunks] + elif file_extension in ('jpg', 'jpeg', 'png', 'webp'): + image = file + _, image_size, image_data = self._get_image_data(file_bytes, file_extension) + else: + raise FileExtensionNotSupported(['PDF', 'DOC', 'DOCX', 'XLSX', 'JPG', 'JPEG', 'PNG', 'WEBP']) + chat_history = self.get_chat_history(model_name=model_name) + chat_history.add_message(HumanMessage(content=input_message.content)) + llm_input = [SystemMessage(content=user_system_prompt), HumanMessage(content=input_content)] + input_tokens, input_embedding_tokens = self._get_input_tokens( + file, image, chunks, chat_history, llm_input, model_name + ) + self.assert_enough_balance( + input_tokens, + image_size, + model=model_name, + embedding_tokens=input_embedding_tokens, + output_tokens=500 if is_free_plan else 4000, + ) + for proxy in Proxy.objects.all(): + system = chat_history.messages.pop(0) + messages = [ + {'role': 'user' if isinstance(msg, HumanMessage) else 'assistant', 'content': msg.content} + for msg in chat_history.messages + ] + messages.insert(0, {'role': 'system', 'content': system.content}) + messages.insert(0, {'role': 'system', 'content': user_system_prompt}) + if image: + messages[-1]['content'] = [ + {'type': 'input_text', 'text': input_message.content}, + {'type': 'input_image', 'image_url': image_data['image_url']['url']}, + ] + elif file: + 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}' + ) + tools = [] + if not is_free_plan: + tools.append( + { + 'type': 'image_generation', + 'size': '1024x1024', + 'quality': 'medium', + 'model': 'gpt-image-1.5', + } + ) + json_data = { + 'model': model_name, + 'input': messages, + 'tools': tools, + 'instructions': ( + 'Форматирование — обязательное требование. Выполняй строго по правилам:\n\n' + "1) Используй реальные символы новой строки, не выводи '\\n' как текст — вставляй переносы.\n\n" + '2) Абзацы: между абзацами ставь две пустые строки (два символа новой строки подряд).\n\n' + '3) Нумерованные и маркированные списки: каждый пункт на отдельной строке;\n' + ' между списком и текстом оставляй две пустые строки.\n\n' + '4) Блоки кода: любые фрагменты кода выделяй тройными бэктиками (```) с указанием языка программирования;\n' + ' перед и после блока оставляй две пустые строки.\n\n' + "5) Заголовки абзацев: делай крупным, используя Markdown '####' (например, '### Заголовок');\n" + ' выделяй жирным (**Заголовок**); оставляй две пустые строки перед и после заголовка.\n\n' + '6) Используй Markdown для всего форматирования, не используй HTML.\n\n' + '7) Исправление формата: если формат неверный, перепиши ответ и верни исправленный вариант.\n\n' + 'Строго разделяй текст на абзацы с жирными заголовками;\n' + 'нумерованные и маркированные списки выводи с переносами строк;\n' + 'блоки кода — с тройными бэктиками и указанием языка;\n' + "не выводи '\\n' как текст, используйте реальные переносы строк;\n" + 'добавляй две пустые строки между абзацами и блоками для улучшения читаемости.' + ), + } + if is_free_plan: + json_data['max_output_tokens'] = 500 + json_data['reasoning'] = {'effort': 'none', 'summary': 'auto'} + elif reasoning := info.get('reasoning'): + reasoning_data = { + 'Минимальный': 'minimal', + 'Низкий': 'low', + 'Средний': 'medium', + 'Высокий': 'high', + 'Сверхвысокий': 'xhigh', + } + json_data['reasoning'] = {'effort': reasoning_data[reasoning], 'summary': 'auto'} + if reasoning == 'Минимальный': + info.pop('web_search', None) + info.pop('code_interpreter', None) + if model_name == 'gpt-5.4' and (verbosity := info.get('verbosity', 'Отключено')) != 'Отключено': + verbosity_data = { + 'Низкий': 'low', + 'Средний': 'medium', + 'Высокий': 'high', + } + json_data['text'] = {'verbosity': verbosity_data[verbosity]} + if (web_search := info.get('web_search', 'Отключено')) != 'Отключено': + search_context_sizes = { + 'Малый контекст': 'low', + 'Средний контекст': 'medium', + 'Большой контекст': 'high', + } + json_data['tools'].append( + { + 'type': 'web_search_preview', + 'search_context_size': search_context_sizes[web_search], + 'user_location': {'type': 'approximate', 'country': 'RU'}, + } + ) + info['web_search'] = search_context_sizes[web_search] + if info.get('code_interpreter') and model_name == 'gpt-5.4': + json_data['tools'].append({'type': 'code_interpreter', 'container': {'type': 'auto'}}) + messages[-1]['content'] += ' the python tool ' + input_tokens, output_tokens, response = self.call_openai_api( + proxy=proxy, endpoint='responses', json_data=json_data + ) + generated_image = None + if isinstance(response.content, list): + if isinstance(response.content[0], dict) and response.content[0].get('generate_image'): + generated_image = base64.b64decode(response.content[0]['image']) + 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}') + 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}' + ) + process_time = timedelta(seconds=time.time() - start_time) + self.handle_invoice( + self.neuron_model, input_tokens, output_tokens, model_name, info, embedding_tokens, generated_image + ) + msgs = self.save_results([response], process_time, generated_image, save) + return msgs @@ -45,7 +45,7 @@ class Claude(SimpleService): }, # 1M tokens } - TOOLS_TOKEN_COSTS = {'text-embedding-3-large': {'output': Decimal('0.000065')}} + TOOLS_TOKEN_COSTS = {'text-embedding-3-small': {'output': Decimal('0.00001')}} def calculate_price( self, version: str, input_tokens: int, output_tokens: int, embedding_tokens: int @@ -55,7 +55,7 @@ class Claude(SimpleService): input_tokens * price_map['input'] / 1_000_000 + output_tokens * price_map['output'] / 1_000_000 ) if embedding_tokens > 0: - price += self.TOOLS_TOKEN_COSTS['text-embedding-3-large']['output'] * embedding_tokens + price += self.TOOLS_TOKEN_COSTS['text-embedding-3-small']['output'] * embedding_tokens return price.quantize(Decimal('0.1'), rounding='ROUND_UP') def save_results(self, content: Iterator[Any], t: timedelta, save: bool = True) -> list[Message]: @@ -98,7 +98,12 @@ class Claude(SimpleService): 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 + 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, @@ -0,0 +1,63 @@ +import time +from datetime import timedelta +from decimal import Decimal +from io import BytesIO +from typing import Any + +import requests +from django.core.files import File +from replicate.exceptions import ModelError + +from messages.models import Message +from ml_model.exceptions import RequestBlocked, GenerationException +from ml_model.services.base import SimpleService +from ml_model.tasks import replicate_run + +from payments.exceptions.insufficient_balance import InsufficientBalance +from payments.selectors.payment_plan_selector import PaymentPlanSelector + + +class Elevenlabs(SimpleService): + TOKENS_COST = Decimal('2.490') + + @classmethod + def predict_price(cls, content: str, file_exists: bool, info: dict[str, Any]) -> Decimal | None: + duration = info['duration'] + return (cls.TOKENS_COST * duration).quantize(Decimal('0.1'), rounding='ROUND_UP') + + def calculate_price(self, duration: int) -> Decimal: + return (self.TOKENS_COST * duration).quantize(Decimal('0.1'), rounding='ROUND_UP') + + def save_results(self, content: str, t: timedelta, video: str, save: bool = True) -> list[Message]: + msg = Message( + content=content, + content_object=self.store, + elapsed_time=t, + file=File(BytesIO(requests.get(video).content), '.mp3'), + ) + if save: + return Message.objects.bulk_create([msg]) + return [msg] + + def make(self, input_message: Message, save: bool = True) -> list[Message]: + duration = input_message.info.pop('duration') + if (balance := PaymentPlanSelector(self.store.user).get_current_balance()) < ( + cost := self.calculate_price(duration) + ): + raise InsufficientBalance(balance, cost) + callback_data = { + 'prompt': input_message.content, + 'music_length_ms': duration * 1000, + **input_message.info, + } + start_time = time.time() + try: + audio = replicate_run('elevenlabs/music', callback_data) + except ModelError as exc: + if any(error in str(exc) for error in ('E005', 'E006', 'sexual')): + raise RequestBlocked + raise GenerationException from exc + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice(input_message.content_object.model, duration=duration) + msgs = self.save_results(input_message.content, process_time, audio, save) + return msgs @@ -9,6 +9,7 @@ from django.db.models.fields.files import FieldFile from PIL import Image from messages.models import Message +from ml_model.exceptions import FileExtensionNotSupported from ml_model.services.EmbeddingService import EmbeddingService from ml_model.services.FileService import FileProcessingService from ml_model.services.base import SimpleService @@ -50,12 +51,6 @@ class Gemini(SimpleService): 'input': Decimal('30'), 'output': Decimal('120'), }, - 'gemini-3-pro-preview': { - 'input': Decimal('600'), - 'output': Decimal('3600'), - 'input_imgs': Decimal('0'), - 'highest_prices': {'input': Decimal('1200'), 'output': Decimal('5400')}, - }, 'gemini-3-flash-preview': { 'input': Decimal('150'), 'output': Decimal('900'), @@ -63,13 +58,13 @@ class Gemini(SimpleService): }, } - TOOLS_TOKEN_COSTS = {'text-embedding-3-large': {'output': Decimal('0.000065')}} + TOOLS_TOKEN_COSTS = {'text-embedding-3-small': {'output': Decimal('0.00001')}} def calculate_price( self, version: str, input_tokens: int, output_tokens: int, image: FieldFile, embedding_tokens: int ) -> Decimal: price_map = self.TOKENS_COST[version.split('/')[1]] - if version.split('/')[1] in ('gemini-2.5-pro', 'gemini-3-pro-preview') and input_tokens > 200_000: + if version.split('/')[1] == 'gemini-2.5-pro' and input_tokens > 200_000: price = ( input_tokens * price_map['highest_prices']['input'] / 1_000_000 + output_tokens * price_map['highest_prices']['output'] / 1_000_000 @@ -82,7 +77,7 @@ class Gemini(SimpleService): if image: price += price_map['input_imgs'] / 1_000 if embedding_tokens > 0: - price += self.TOOLS_TOKEN_COSTS['text-embedding-3-large']['output'] * embedding_tokens + price += self.TOOLS_TOKEN_COSTS['text-embedding-3-small']['output'] * embedding_tokens return price.quantize(Decimal('0.1'), rounding='ROUND_UP') def save_results(self, content: str, t: timedelta, save: bool = True) -> list[Message]: @@ -134,10 +129,7 @@ class Gemini(SimpleService): 'Priority: analytical depth, internal consistency, and correctness over speed.', }, ) - callback_data.update({ - "reasoning": {"effort": "high"}, - "temperature": 0.2 - }) + callback_data.update({'reasoning': {'effort': 'high'}, 'temperature': 0.2}) file = input_message.file image = None embedding_tokens = 0 @@ -152,11 +144,14 @@ class Gemini(SimpleService): 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] - ) + 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 + 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, @@ -168,8 +163,7 @@ class Gemini(SimpleService): f'Используй системный промпт. Содержание файла: ' f'{chunks}. Вопрос: {input_message.content}' ) - else: - kind = filetype.guess(file_bytes[:20]) + elif file_extension in ('jpg', 'jpeg', 'png', 'webp'): mime = kind.mime if kind else 'application/octet-stream' normalized_image = Image.open(file) format = 'jpeg' if kind.extension == 'jpg' else kind.extension @@ -182,6 +176,8 @@ class Gemini(SimpleService): {'type': 'image_url', 'image_url': {'url': image_url}}, ] image = file + else: + raise FileExtensionNotSupported(['PDF', 'DOC', 'DOCX', 'XLSX', 'JPG', 'JPEG', 'PNG', 'WEBP']) start_time = time.time() result = openrouter_run(version, messages, callback_data, 'Gemini') process_time = timedelta(seconds=(time.time() - start_time)) @@ -0,0 +1,179 @@ +import base64 +import time +from datetime import timedelta +from decimal import Decimal +from io import BytesIO + +import filetype +from PIL import Image + +from messages.models import Message +from ml_model.exceptions import FileExtensionNotSupported +from ml_model.services.EmbeddingService import EmbeddingService +from ml_model.services.FileService import FileProcessingService +from ml_model.services.base import SimpleService +from ml_model.tasks import openrouter_run +from poller.models import Proxy +from tools.chats.models import Chat +from tools.copywrite.models import Copywrite +from tools.public_api.models import APIStore + + +class Gemini_3_1(SimpleService): + TOKENS_COST = { + 'gemini-3.1-pro-preview': { + 'input': Decimal('600'), + 'output': Decimal('3600'), + 'highest_prices': {'input': Decimal('1200'), 'output': Decimal('5400')}, + }, + 'gemini-3.1-flash-lite-preview': { + 'input': Decimal('75'), + 'output': Decimal('450'), + 'highest_prices': {'input': Decimal('75'), 'output': Decimal('450')}, + } + } + + TOOLS_TOKEN_COSTS = {'text-embedding-3-small': {'output': Decimal('0.00001')}} + + 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: + price = ( + input_tokens * price_map['highest_prices']['input'] / 1_000_000 + + output_tokens * price_map['highest_prices']['output'] / 1_000_000 + ) + else: + price = ( + input_tokens * price_map['input'] / 1_000_000 + + output_tokens * price_map['output'] / 1_000_000 + ) + if embedding_tokens > 0: + price += self.TOOLS_TOKEN_COSTS['text-embedding-3-small']['output'] * embedding_tokens + price += Decimal('2') + return price.quantize(Decimal('0.1'), rounding='ROUND_UP') + + def save_results(self, content: str, t: timedelta, save: bool = True) -> list[Message]: + msgs = [ + Message( + content=content, + content_object=self.store, + elapsed_time=t, + ) + ] + if save: + return Message.objects.bulk_create(msgs) + return msgs + + def make(self, input_message: Message, save: bool = True) -> list[Message]: + version = input_message.info.get('version', 'gemini-3.1-pro-preview:online') + callback_data = { + 'provider': {'order': ['Google AI Studio']}, + **input_message.info, + } + messages = self.get_chat_history() + messages.insert( + 0, + { + 'role': 'system', + 'content': ( + "Always respond in the same language as the user's last message, " + 'unless the user explicitly asks you to answer in a different language.' + ), + }, + ) + messages.append({'role': 'user', 'content': input_message.content}) + embedding_tokens = 0 + if input_message.file: + file_service = FileProcessingService + file_bytes = input_message.file.read() + kind = filetype.guess(file_bytes[:20]) + 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' + normalized_image = Image.open(input_message.file) + format = 'jpeg' if kind.extension == 'jpg' else kind.extension + buf = BytesIO() + normalized_image.save(buf, format=format) + image_url = f'data:{mime};base64,{base64.b64encode(buf.getvalue()).decode("utf-8")}' + buf.close() + messages[-1]['content'] = [ + {'type': 'text', 'text': input_message.content}, + {'type': 'image_url', 'image_url': {'url': image_url}}, + ] + else: + raise FileExtensionNotSupported(['PDF', 'DOC', 'DOCX', 'XLSX', 'JPG', 'JPEG', 'PNG', 'WEBP']) + start_time = time.time() + result = openrouter_run(f'google/{version}:online', messages, callback_data, 'Gemini 3.1') + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice( + input_message.content_object.model, + version=version, + input_tokens=result[1], + output_tokens=result[2], + embedding_tokens=embedding_tokens, + ) + msgs = self.save_results(result[0], 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] + ) + ) + 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 @@ -9,7 +9,7 @@ from django.core.files import File from replicate.exceptions import ModelError from messages.models import Message -from ml_model.exceptions import ModelTimeoutError, ImageContentNotFound +from ml_model.exceptions import ModelTimeoutError, ImageContentNotFound, RequestBlocked, GenerationException from ml_model.services.base import SimpleService from ml_model.tasks import replicate_run @@ -57,6 +57,12 @@ class Geminiimage(SimpleService): msgs = self.save_results(input_message.content, process_time, images, save) return msgs except ModelError as exc: - if exc.prediction.error == 'No image content found in response': + if exc.prediction.error in ( + 'No image content found in response', + 'Failed to generate image.', + ): raise ImageContentNotFound + elif any(error in str(exc) for error in ('E005', 'E006', 'sexual')): + raise RequestBlocked + raise GenerationException from exc raise ModelTimeoutError @@ -2,17 +2,13 @@ import base64 import time from datetime import timedelta from decimal import Decimal -from io import BytesIO import filetype -from PIL import Image from messages.models import Message -from ml_model.services.EmbeddingService import EmbeddingService -from ml_model.services.FileService import FileProcessingService +from ml_model.exceptions import CorruptedFileError, FileExtensionNotSupported from ml_model.services.base import SimpleService from ml_model.tasks import openrouter_run -from poller.models import Proxy from tools.chats.models import Chat from tools.copywrite.models import Copywrite from tools.public_api.models import APIStore @@ -52,8 +48,12 @@ class Gemma(SimpleService): messages.append({'role': 'user', 'content': input_message.content}) if input_message.file: kind = filetype.guess(input_message.file.read(20)) + if not kind: + raise CorruptedFileError mime = kind.mime if kind else 'application/octet-stream' input_message.file.seek(0) + if kind.extension.upper() not in (available_extensions := ('JPG', 'JPEG', 'PNG', 'WEBP')): + raise FileExtensionNotSupported(available_extensions) image_url = f'data:{mime};base64,{base64.b64encode(input_message.file.read()).decode("utf-8")}' input_message.file.close() messages[-1]['content'] = [ @@ -1,14 +1,22 @@ import base64 +import logging import time from datetime import timedelta from decimal import Decimal -from typing import Any, Iterator +from io import BytesIO import filetype +from PIL import Image from messages.models import Message +from ml_model.exceptions import FileExtensionNotSupported, CorruptedFileError +from ml_model.services.EmbeddingService import EmbeddingService +from ml_model.services.FileService import FileProcessingService from ml_model.services.base import SimpleService from ml_model.tasks import openrouter_run +from payments.exceptions.insufficient_balance import InsufficientBalance +from payments.selectors.payment_plan_selector import PaymentPlanSelector +from poller.models import Proxy from tools.chats.models import Chat from tools.copywrite.models import Copywrite from tools.public_api.models import APIStore @@ -17,14 +25,18 @@ from tools.public_api.models import APIStore class Grok_4_1_Fast(SimpleService): TOKENS_COST = {'input': Decimal('40'), 'output': Decimal('100')} - def calculate_price(self, input_tokens: int, output_tokens: int) -> Decimal: + TOOLS_TOKEN_COSTS = {'text-embedding-3-small': {'output': Decimal('0.00001')}} + + def calculate_price(self, input_tokens: int, output_tokens: int, embedding_tokens: int) -> Decimal: price = ( input_tokens * self.TOKENS_COST['input'] / 1_000_000 + output_tokens * self.TOKENS_COST['output'] / 1_000_000 ) + if embedding_tokens > 0: + price += self.TOOLS_TOKEN_COSTS['text-embedding-3-small']['output'] * embedding_tokens return price.quantize(Decimal('0.01'), 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,16 +55,67 @@ class Grok_4_1_Fast(SimpleService): } messages = self.get_chat_history() messages.append({'role': 'user', 'content': input_message.content}) + embedding_tokens = 0 if input_message.file: - kind = filetype.guess(input_message.file.read(20)) - mime = kind.mime if kind else 'application/octet-stream' - input_message.file.seek(0) - image_url = f'data:{mime};base64,{base64.b64encode(input_message.file.read()).decode("utf-8")}' - input_message.file.close() - messages[-1]['content'] = [ - {'type': 'text', 'text': input_message.content}, - {'type': 'image_url', 'image_url': {'url': image_url}}, - ] + file_service = FileProcessingService + file_bytes = input_message.file.read() + kind = filetype.guess(file_bytes[:20]) + if not kind: + 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) + approx_tokens = sum([len(message['content']) for message in messages]) / 3 + predict_price = ( + Decimal(approx_tokens) * self.TOKENS_COST['input'] / Decimal('1000000') + + len(chunks) * 2100 * self.TOOLS_TOKEN_COSTS['text-embedding-3-small']['output'] + ).quantize(Decimal('0.1'), rounding='ROUND_UP') + if (balance := PaymentPlanSelector(self.store.user).get_current_balance()) < predict_price: + if self.store.user.payment_plan.plan.price <= 0: + return self.save_results( + content='Файл не удаётся обработать — его размер больше максимально допустимого ' + 'для вашего тарифа. Для продолжения выберите план с увеличенным лимитом.', + t=timedelta(minutes=0, seconds=0), + ) + raise InsufficientBalance(balance, predict_price) + 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' + normalized_image = Image.open(input_message.file) + format = 'jpeg' if kind.extension == 'jpg' else kind.extension + buf = BytesIO() + normalized_image.save(buf, format=format) + image_url = f'data:{mime};base64,{base64.b64encode(buf.getvalue()).decode("utf-8")}' + buf.close() + messages[-1]['content'] = [ + {'type': 'text', 'text': input_message.content}, + {'type': 'image_url', 'image_url': {'url': image_url}}, + ] + else: + raise FileExtensionNotSupported(['PDF', 'DOC', 'DOCX', 'XLSX', 'JPG', 'JPEG', 'PNG', 'WEBP']) start_time = time.time() result = openrouter_run('x-ai/grok-4.1-fast', messages, callback_data, 'Grok 4.1 Fast') process_time = timedelta(seconds=(time.time() - start_time)) @@ -60,6 +123,7 @@ class Grok_4_1_Fast(SimpleService): input_message.content_object.model, input_tokens=result[1], output_tokens=result[2], + embedding_tokens=embedding_tokens, ) msgs = self.save_results(result[0], process_time) return msgs @@ -0,0 +1,64 @@ +import base64 +import time +from datetime import timedelta +from decimal import Decimal +from io import BytesIO +from typing import Any + +import filetype +import requests +from django.core.files import File +from replicate.exceptions import ModelError + +from messages.models import Message +from ml_model.exceptions import RequestBlocked, GenerationException +from ml_model.services.base import SimpleService +from ml_model.tasks import replicate_run + +from payments.exceptions.insufficient_balance import InsufficientBalance +from payments.selectors.payment_plan_selector import PaymentPlanSelector + + +class Grok_Imagine_Video(SimpleService): + TOKENS_COST = Decimal('15') + + @classmethod + def predict_price(cls, content: str, file_exists: bool, info: dict[str, Any]) -> Decimal | None: + duration = info['duration'] + price = cls.TOKENS_COST * duration + return price.quantize(Decimal('0.1'), rounding='ROUND_UP') + + def calculate_price(self, duration: int) -> Decimal: + return (self.TOKENS_COST * duration).quantize(Decimal('0.1'), rounding='ROUND_UP') + + def save_results(self, content: str, t: timedelta, video: str, save: bool = True) -> list[Message]: + msg = Message( + content=content, + content_object=self.store, + elapsed_time=t, + file=File(BytesIO(requests.get(video).content), '.mp4'), + ) + if save: + return Message.objects.bulk_create([msg]) + return [msg] + + def make(self, input_message: Message, save: bool = True) -> list[Message]: + duration = input_message.info.get('duration', 5) + if (balance := PaymentPlanSelector(self.store.user).get_current_balance()) < ( + cost := self.TOKENS_COST * duration + ): + raise InsufficientBalance(balance, cost) + callback_data = dict({'prompt': self.translate_prompt(input_message.content), **input_message.info}) + if image := input_message.file: + callback_data.update({'image': image.url}) + start_time = time.time() + try: + video = replicate_run('xai/grok-imagine-video', callback_data) + except ModelError as exc: + if any(error in str(exc) for error in ('E005', 'E006', 'sexual')): + raise RequestBlocked + raise GenerationException from exc + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice(input_message.content_object.model, duration=duration) + msgs = self.save_results(input_message.content, process_time, video, save) + return msgs @@ -59,4 +59,4 @@ class Imagen(SimpleService): raise RequestBlocked elif exc.prediction.error == 'No image content found in response': raise ImageContentNotFound - raise GenerationException + raise GenerationException from exc @@ -0,0 +1,79 @@ +import time +from datetime import timedelta +from decimal import Decimal +from io import BytesIO +from typing import Any + +import requests +from django.core.files import File + +from messages.models import Message +from ml_model.services.base import SimpleService +from ml_model.tasks import replicate_run + +from payments.exceptions.insufficient_balance import InsufficientBalance +from payments.selectors.payment_plan_selector import PaymentPlanSelector + + +class Ltx(SimpleService): + TOKENS_COST = { + '1080p': Decimal('12'), + '2k': Decimal('24'), + '4k': Decimal('48'), + } + + @classmethod + def predict_price(cls, content: str, file_exists: bool, info: dict[str, Any]) -> Decimal | None: + resolution = info['resolution'] + duration = info['duration'] + return (cls.TOKENS_COST[resolution] * duration).quantize(Decimal('0.1'), rounding='ROUND_UP') + + def calculate_price(self, resolution: str, duration: int) -> Decimal: + return (self.TOKENS_COST[resolution] * duration).quantize(Decimal('0.1'), rounding='ROUND_UP') + + def save_results(self, content: str, t: timedelta, video: str, save: bool = True) -> list[Message]: + msg = Message( + content=content, + content_object=self.store, + elapsed_time=t, + file=File(BytesIO(requests.get(video).content), '.mp4'), + ) + if save: + return Message.objects.bulk_create([msg]) + return [msg] + + def make(self, input_message: Message, save: bool = True) -> list[Message]: + resolution = input_message.info.pop('resolution', '1080p') + duration = input_message.info.pop('duration', 6) + if (balance := PaymentPlanSelector(self.store.user).get_current_balance()) < ( + cost := self.TOKENS_COST[resolution] * duration + ): + raise InsufficientBalance(balance, cost) + camera_motion = { + 'Без движения камеры': 'none', + 'Приближение камеры': 'dolly_in', + 'Удаление камеры': 'dolly_out', + 'Движение камеры влево': 'dolly_left', + 'Движение камеры вправо': 'dolly_right', + 'Подъём камеры': 'jib_up', + 'Опускание камеры': 'jib_down', + 'Статичная камера': 'static', + 'Смена фокуса': 'focus_shift', + } + callback_data = dict( + { + 'prompt': input_message.content, + 'resolution': resolution, + 'duration': duration, + 'camera_motion': camera_motion[input_message.info.pop('camera_motion', 'Без движения камеры')], + **input_message.info, + } + ) + if input_message.file: + callback_data.update({'image': input_message.file.url}) + start_time = time.time() + video = replicate_run('lightricks/ltx-2.3-fast', callback_data) + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice(input_message.content_object.model, resolution=resolution, duration=duration) + msgs = self.save_results(input_message.content, process_time, video, save) + return msgs @@ -0,0 +1,79 @@ +import time +from datetime import timedelta +from decimal import Decimal +from io import BytesIO +from typing import Optional, Any + +import requests +from django.core.files import File +from replicate.exceptions import ModelError + +from messages.models import Message +from ml_model.exceptions import ( + ImageContentNotFound, + GenerationException, + RequestBlocked, + ServiceHighDemandError, +) +from ml_model.services.base import SimpleService +from ml_model.tasks import replicate_run + + +class Nanobanana_2(SimpleService): + TOKENS_COST = { + '1K': Decimal('20.1'), + '2K': Decimal('30.3'), + '4K': Decimal('45.3'), + } + + def calculate_price(self, resolution: Optional[str]) -> Decimal: + return self.TOKENS_COST[resolution] + + @classmethod + def predict_price(cls, content: str, file_exists: bool, info: dict[str, Any]) -> Decimal | None: + return cls.TOKENS_COST[info['resolution']] + + def save_results( + self, + prompt: str, + image_url: str, + time: timedelta, + save: bool = True, + ) -> list[Message]: + message = Message( + content_object=self.store, + elapsed_time=time, + content=prompt, + file=File(BytesIO(requests.get(image_url).content), '.png'), + ) + if save: + return Message.objects.bulk_create([message]) + return [message] + + def make(self, input_message: Message, save: bool = True) -> list[Message]: + start_time = time.time() + resolution = input_message.info.get('resolution', '2K') + callback_data = dict( + { + 'prompt': self.translate_prompt(input_message.content), + **input_message.info, + } + ) + if input_message.file: + callback_data.update( + {'image_input': [input_message.file.url], 'aspect_ratio': 'match_input_image'} + ) + try: + image = replicate_run('google/nano-banana-2', callback_data) + except ModelError as exc: + if exc.prediction.error == 'No image content found in response': + raise ImageContentNotFound + elif any(error in str(exc) for error in ('E005', 'E006', 'sexual')): + raise RequestBlocked + elif 'E003' in str(exc): + raise ServiceHighDemandError from exc + raise GenerationException from exc + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice(input_message.content_object.model, resolution=resolution) + msgs = self.save_results(input_message.content, image, process_time, save) + return msgs @@ -13,7 +13,12 @@ from django.core.files import File from replicate.exceptions import ModelError from messages.models import Message -from ml_model.exceptions import ImageContentNotFound, GenerationException, RequestBlocked +from ml_model.exceptions import ( + ImageContentNotFound, + GenerationException, + RequestBlocked, + FileTooLargeError, +) from ml_model.services.base import SimpleService from ml_model.tasks import replicate_run from payments.exceptions.insufficient_balance import InsufficientBalance @@ -51,6 +56,8 @@ class Photon(SimpleService): } ) if input_message.file: + if input_message.file.size >= (10 << 10 << 10): + raise FileTooLargeError(10) kind = filetype.guess(input_message.file.read(20)) mime = kind.mime if kind else 'application/octet-stream' input_message.file.seek(0) @@ -65,7 +72,7 @@ class Photon(SimpleService): raise RequestBlocked elif exc.prediction.error == 'No image content found in response': raise ImageContentNotFound - raise GenerationException + raise GenerationException from exc process_time = timedelta(seconds=(time.time() - start_time)) self.handle_invoice(input_message.content_object.model) msgs = self.save_results(input_message.content, process_time, images, save) @@ -0,0 +1,93 @@ +import time + +from datetime import timedelta +from decimal import Decimal +from io import BytesIO +from typing import Any + +import filetype +import requests + +from django.core.files import File +from replicate.exceptions import ModelError + +from messages.models import Message +from ml_model.exceptions import GenerationException, RequestBlocked +from ml_model.services.FileService import FileProcessingService +from ml_model.services.base import SimpleService +from ml_model.tasks import replicate_run +from payments.exceptions.insufficient_balance import InsufficientBalance +from payments.selectors.payment_plan_selector import PaymentPlanSelector + + +class Pruna_V(SimpleService): + TOKENS_COST = { + '720p': {'standard': Decimal('10'), 'draft': Decimal('2.5')}, + '1080p': {'standard': Decimal('20'), 'draft': Decimal('5')}, + } + + def calculate_price(self, resolution: str, generation_mode: str, duration: int) -> Decimal: + return self.TOKENS_COST[resolution][generation_mode] * duration + + @classmethod + def predict_price(cls, content: str, file_exists: bool, info: dict[str, Any]) -> Decimal | None: + resolution = info['resolution'] + generation_mode = info['generation_mode'] + duration = info['duration'] + return cls.TOKENS_COST[resolution][generation_mode] * duration + + def save_results(self, content: str, t: timedelta, image_url: str, save: bool = True) -> list[Message]: + msg = Message( + content=content, + content_object=self.store, + elapsed_time=t, + file=File(BytesIO(requests.get(image_url).content), '.mp4'), + ) + if save: + return Message.objects.bulk_create([msg]) + return [msg] + + def make(self, input_message: Message, save: bool = True) -> list[Message]: + resolution = input_message.info.get('resolution', '720p') + generation_mode = input_message.info.pop('generation_mode', 'standard') + duration = input_message.info.get('duration', 5) + if (balance := PaymentPlanSelector(self.store.user).get_current_balance()) < ( + cost := self.TOKENS_COST[resolution][generation_mode] * duration + ): + raise InsufficientBalance(balance, cost) + callback_data = dict( + { + 'prompt': self.translate_prompt(input_message.content), + 'draft': generation_mode == 'draft', + 'disable_safety_filter': False, + **input_message.info, + } + ) + if input_message.file: + file_service = FileProcessingService + file_bytes = input_message.file.read() + kind = filetype.guess(file_bytes[:20]) + raw_file_extension = kind.extension + file_extension = file_service.get_file_extension(raw_file_extension, file_bytes) + if file_extension in ('flac', 'mp3', 'wav'): + callback_data.update({'audio': input_message.file.url}) + else: + callback_data.update({'image': input_message.file.url}) + start_time = time.time() + try: + images = replicate_run('prunaai/p-video', callback_data) + if 'nsfw.jpeg' == str(images).split('/')[-1]: + raise RequestBlocked + except ModelError as exc: + if any(error in str(exc) for error in ('E005', 'E006', 'sexual', 'NSFW')): + raise RequestBlocked + raise GenerationException from exc + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice( + input_message.content_object.model, + resolution=resolution, + generation_mode=generation_mode, + duration=duration, + ) + msgs = self.save_results(input_message.content, process_time, images, save) + return msgs @@ -0,0 +1,82 @@ +import time + +from datetime import timedelta +from decimal import Decimal +from io import BytesIO +from typing import Any + +import requests + +from django.core.files import File +from replicate.exceptions import ModelError + +from messages.models import Message +from ml_model.exceptions import GenerationException, RequestBlocked +from ml_model.services.base import SimpleService +from ml_model.tasks import replicate_run +from payments.exceptions.insufficient_balance import InsufficientBalance +from payments.selectors.payment_plan_selector import PaymentPlanSelector + + +class Prunaai(SimpleService): + TOKENS_COST = {'p-image': Decimal('2.5'), 'p-image-edit': Decimal('5'), 'flux-fast': Decimal('2.5')} + + def calculate_price(self, version: str) -> Decimal: + return self.TOKENS_COST[version] + + @classmethod + def predict_price(cls, content: str, file_exists: bool, info: dict[str, Any]) -> Decimal | None: + version = info['version'] + if file_exists: + version = 'p-image-edit' + return cls.TOKENS_COST[version] + + def save_results(self, content: str, t: timedelta, image_url: str, save: bool = True) -> list[Message]: + msg = Message( + content=content, + content_object=self.store, + elapsed_time=t, + file=File(BytesIO(requests.get(image_url).content), '.png'), + ) + if save: + return Message.objects.bulk_create([msg]) + return [msg] + + def make(self, input_message: Message, save: bool = True) -> list[Message]: + version = input_message.info.get('version', 'p-image') + aspect_ratio = input_message.info.pop('aspect_ratio', 'custom') + if input_message.file: + version = 'p-image-edit' + aspect_ratio = 'match_input_image' + if (balance := PaymentPlanSelector(self.store.user).get_current_balance()) < self.TOKENS_COST[ + version + ]: + raise InsufficientBalance(balance, self.TOKENS_COST[version]) + callback_data = dict( + { + 'prompt': self.translate_prompt(input_message.content), + 'aspect_ratio': aspect_ratio, + **input_message.info, + } + ) + if input_message.file: + callback_data.update({'images': [input_message.file.url]}) + if version == 'flux-fast' and (s_m := input_message.info.get('speed_mode', None)): + speed_mode = { + 'Легкий сок 🍊 (более стабильный результат)': 'Lightly Juiced 🍊 (more consistent)', + 'Сок 🔥 (режим по умолчанию)': 'Juiced 🔥 (default)', + 'Экстра-сок 🔥 (быстрее генерация)': 'Extra Juiced 🔥 (more speed)', + 'Мгновение ока 👁️ (максимальная скорость)': 'Blink of an eye 👁️', + } + callback_data.update({'speed_mode': speed_mode[s_m]}) + start_time = time.time() + try: + images = replicate_run(f'prunaai/{version}', callback_data) + except ModelError as exc: + if any(error in str(exc) for error in ('E005', 'E006', 'sexual', 'NSFW')): + raise RequestBlocked + raise GenerationException from exc + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice(input_message.content_object.model, version=version) + msgs = self.save_results(input_message.content, process_time, images, save) + return msgs @@ -0,0 +1,166 @@ +import base64 +import time +from datetime import timedelta +from decimal import Decimal +from io import BytesIO + +import filetype +from PIL import Image + +from messages.models import Message +from ml_model.exceptions import FileExtensionNotSupported +from ml_model.services.EmbeddingService import EmbeddingService +from ml_model.services.FileService import FileProcessingService +from ml_model.services.base import SimpleService +from ml_model.tasks import openrouter_run +from poller.models import Proxy +from tools.chats.models import Chat +from tools.copywrite.models import Copywrite +from tools.public_api.models import APIStore + + +class Qwen_3_5(SimpleService): + TOKENS_COST = { + 'qwen3.5-9b': {'input': Decimal('30'), 'output': Decimal('45')}, + 'qwen3.5-flash-02-23': {'input': Decimal('30'), 'output': Decimal('120')}, + 'qwen3.5-35b-a3b': {'input': Decimal('48.75'), 'output': Decimal('390')}, + 'qwen3.5-27b': {'input': Decimal('58.5'), 'output': Decimal('468')}, + 'qwen3.5-122b-a10b': {'input': Decimal('78'), 'output': Decimal('624')}, + 'qwen3.5-397b-a17b': {'input': Decimal('117'), 'output': Decimal('702')}, + } + + PROVIDERS = { + 'qwen3.5-9b': 'together', + 'qwen3.5-flash-02-23': 'alibaba', + 'qwen3.5-35b-a3b': 'alibaba', + 'qwen3.5-27b': 'alibaba', + 'qwen3.5-122b-a10b': 'alibaba', + 'qwen3.5-397b-a17b': 'alibaba', + } + + TOOLS_TOKEN_COSTS = {'text-embedding-3-small': {'output': Decimal('0.00001')}} + + def calculate_price( + self, version: str, input_tokens: int, output_tokens: int, embedding_tokens: int + ) -> Decimal: + price = ( + input_tokens * self.TOKENS_COST[version]['input'] / 1_000_000 + + output_tokens * self.TOKENS_COST[version]['output'] / 1_000_000 + + Decimal('2') + ) + if embedding_tokens > 0: + price += self.TOOLS_TOKEN_COSTS['text-embedding-3-small']['output'] * embedding_tokens + return price.quantize(Decimal('0.1'), rounding='ROUND_UP') + + def save_results(self, content: str, time: timedelta, save: bool = True) -> list[Message]: + msgs = [ + Message( + content=content, + content_object=self.store, + elapsed_time=time, + ) + ] + if save: + return Message.objects.bulk_create(msgs) + return msgs + + def make(self, input_message: Message, save: bool = True) -> list[Message]: + start_time = time.time() + version = input_message.info.get('version', 'qwen3.5-9b') + callback_data = {'provider': {'order': [self.PROVIDERS[version]]}, **input_message.info} + messages = self.get_chat_history() + messages.append({'role': 'user', 'content': input_message.content}) + embedding_tokens = 0 + if input_message.file: + file_service = FileProcessingService + file_bytes = input_message.file.read() + kind = filetype.guess(file_bytes[:20]) + 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' + normalized_image = Image.open(input_message.file) + format = 'jpeg' if kind.extension == 'jpg' else kind.extension + buf = BytesIO() + normalized_image.save(buf, format=format) + image_url = f'data:{mime};base64,{base64.b64encode(buf.getvalue()).decode("utf-8")}' + buf.close() + messages[-1]['content'] = [ + {'type': 'text', 'text': input_message.content}, + {'type': 'image_url', 'image_url': {'url': image_url}}, + ] + else: + raise FileExtensionNotSupported(['PDF', 'DOC', 'DOCX', 'XLSX', 'JPG', 'JPEG', 'PNG', 'WEBP']) + model_slug = f'qwen/{version}:online' if self.PROVIDERS[version] == 'alibaba' else f'qwen/{version}' + result = openrouter_run(model_slug, messages, callback_data, 'Qwen 3.5') + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice( + input_message.content_object.model, + version=version, + input_tokens=result[1], + output_tokens=result[2], + embedding_tokens=embedding_tokens, + ) + msgs = self.save_results(result[0], 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 \ No newline at end of file @@ -1,9 +1,12 @@ import re +import subprocess import time from concurrent.futures import ThreadPoolExecutor from concurrent.futures._base import as_completed from typing import List, Tuple +import docx2txt +import filetype import httpx import base64 @@ -20,17 +23,20 @@ from django.template.loader import get_template from openai import BadRequestError from backend import settings -from ml_model.exceptions import TemplateNotFound, TemplateUnknownException, FileExtensionNotSupported, \ - ExceededContextLengthError +from ml_model.exceptions import ( + TemplateNotFound, + TemplateUnknownException, + FileExtensionNotSupported, + ExceededContextLengthError, + CorruptedFileError, +) from ml_model.models import NeuronModel from ml_model.services import Chatgpt from django.core.files.uploadedfile import UploadedFile -from django.utils.translation import gettext_lazy as _ from datetime import timedelta -from pathlib import Path from langchain_core.messages import ( HumanMessage, @@ -40,14 +46,18 @@ from langchain_core.runnables import RunnableWithMessageHistory from langchain_openai.chat_models import ChatOpenAI from messages.models import Message -from ml_model.exceptions import GenerationException +from ml_model.services.EmbeddingService import EmbeddingService +from ml_model.services.FileService import FileProcessingService from poller.models import Proxy from ml_model.tasks import drop_redis_vectors from ml_model.constants import ANCHORS + class Raifgpt(Chatgpt): + EMBEDDING_MODEL_FOR_BILLING = 'text-embedding-3-large' + @property def neuron_model(self): return NeuronModel.objects.get(slug='raifgpt') @@ -68,15 +78,22 @@ class Raifgpt(Chatgpt): embedding_tokens = 0 file = input_message.file if file: - file_extension = Path(file.name).suffix - if file_extension == '.pdf': - raw_text = self.get_pdf_data(file) - elif file_extension in ('.doc', '.docx'): - raw_text = self.get_word_data(file_extension[1:], file.read()) + file_service = FileProcessingService + file_bytes = file.read() + kind = filetype.guess(file_bytes[:20]) + if not kind: + raise CorruptedFileError + raw_file_extension = kind.extension + file_extension = file_service.get_file_extension(raw_file_extension, file_bytes) + if file_extension == 'pdf': + raw_text = self.get_pdf_data(file_bytes) + elif file_extension in ('doc', 'docx'): + raw_text = self.get_word_data(file_extension, file_bytes) else: raise FileExtensionNotSupported(['PDF', 'DOC', 'DOCX']) text = re.sub(r'\n{2,}', '\n', raw_text) - chunks = self.split_text_to_chunks(text, chunk_size=1000) + text_chunks = EmbeddingService.split_text_to_chunks(text, chunk_size=1000) + chunks = [HumanMessage(content=chunk_text) for chunk_text in text_chunks] for proxy in Proxy.objects.all(): self.llm = ChatOpenAI( model='gpt-4o', @@ -107,7 +124,9 @@ class Raifgpt(Chatgpt): input_tokens = self.count_text_tokens([*chat_history.messages]) if sum([len(chunk.content) for chunk in chunks]) > 40_000: redis_client = redis.Redis(host=settings.REDIS_HOST, port=settings.REDIS_PORT, db=0) - document_name = chunks[0].content.partition(f':{chr(10)}')[2].split(f'{chr(10)}')[0][:100] + document_name = ( + chunks[0].content.partition(f':{chr(10)}')[2].split(f'{chr(10)}')[0][:100] + ) message_uid = str(self.store.messages.first().pk).replace('-', '_') with httpx.Client( base_url='https://api.openai.com/v1/', @@ -120,7 +139,12 @@ class Raifgpt(Chatgpt): for chunk_id, chunk in enumerate(chunks): threads.append( executor.submit( - self.process_chunk, client, chunk, redis_client, message_uid, chunk_id + EmbeddingService.process_chunk, + client, + chunk.content, + redis_client, + message_uid, + chunk_id, ) ) for thread in as_completed(threads): @@ -130,7 +154,11 @@ class Raifgpt(Chatgpt): anchor_embeddings = {} with ThreadPoolExecutor(max_workers=settings.MAX_THREADS) as executor: for identify, value in ANCHORS.items(): - threads.append(executor.submit(self.get_anchor_embedding, client, value[0], identify)) + threads.append( + executor.submit( + self.get_anchor_embedding, client, value[0], identify + ) + ) for thread in as_completed(threads): thread_result = thread.result() embedding_tokens += thread_result[1] @@ -140,31 +168,41 @@ class Raifgpt(Chatgpt): for identify, embeddings in anchor_embeddings.items(): threads.append( executor.submit( - self.search_via_embeddings, + EmbeddingService.search_via_embeddings, redis_client, message_uid, - embeddings, - top_k=ANCHORS[identify][1] + user_query_embeddings=embeddings, + top_k=ANCHORS[identify][1], ) ) - result = [s['section_text'] for thread in as_completed(threads) for s in thread.result()] + result = [ + s['section_text'] + for thread in as_completed(threads) + for s in thread.result() + ] else: - query_embedding, e_total_tokens = self.get_embedding(client=client, content=input_message.content) + query_embedding, e_total_tokens = EmbeddingService._get_embedding( + client=client, content=input_message.content + ) embedding_tokens += e_total_tokens result = [ s['section_text'] - for s in self.search_via_embeddings( + for s in EmbeddingService.search_via_embeddings( redis_client=redis_client, message_uid=message_uid, user_query_embeddings=query_embedding, - top_k=25 + top_k=25, ) ] user_input = [ SystemMessage(content=user_system_prompt), - HumanMessage(self.make_embeddings_prompt( - document_name=document_name, section_texts=result, question=input_message.content - )) + HumanMessage( + self.make_embeddings_prompt( + document_name=document_name, + section_texts=result, + question=input_message.content, + ) + ), ] input_tokens += self.count_text_tokens(user_input) response = conversation.invoke( @@ -177,9 +215,11 @@ class Raifgpt(Chatgpt): input = [ SystemMessage(content=user_system_prompt), HumanMessage( - content=f'Используй системный промпт. Содержание файла: ' - f'{chunks}. Вопрос: {input_message.content}' + content=( + 'Используй системный промпт. Содержание файла: ' + f'{"".join(chunk.content for chunk in chunks)}. Вопрос: {input_message.content}' ) + ), ] input_tokens += self.count_text_tokens(input) response = conversation.invoke( @@ -202,23 +242,19 @@ class Raifgpt(Chatgpt): self.logger.info(f'Input количество токенов для raifgpt - {input_tokens}') self.logger.info(f'Output количество токенов для raifgpt - {output_tokens}') self.logger.info(f'Embedding количество токенов для raifgpt - {embedding_tokens}') - self.logger.info(f'Общее количество токенов для raifgpt - {input_tokens + output_tokens + embedding_tokens}') + self.logger.info( + f'Общее количество токенов для raifgpt - {input_tokens + output_tokens + embedding_tokens}' + ) process_time = timedelta(seconds=time.time() - start_time) self.handle_invoice( - self.neuron_model, - input_tokens, - output_tokens, - self.llm.model_name, - {}, - embedding_tokens + self.neuron_model, input_tokens, output_tokens, self.llm.model_name, {}, embedding_tokens ) msgs = self.save_results([response], process_time, save) return msgs - def get_pdf_data(self, pdf_file: UploadedFile) -> str: + def get_pdf_data(self, pdf_data: bytes) -> str: max_batch_size = 3.9 * 1024 * 1024 - pdf_data = pdf_file.read() image_count = 0 try: doc = fitz.open(stream=pdf_data, filetype="pdf") @@ -346,20 +382,43 @@ class Raifgpt(Chatgpt): """ def get_anchor_embedding(self, client: httpx.Client, content: str, anchor: str) -> Tuple[List[float], int, str]: - ''' + """ A method for converting raw text (anchor content) into embeddings using OpenAI API request :param client: Httpx client :param content: raw text of a chunk :param anchor: anchor identifier - ''' - response = client.post( - url="embeddings", - json={ - 'model': 'text-embedding-3-large', - 'input': content - } - ) + """ + response = client.post(url='embeddings', json={'model': 'text-embedding-3-large', 'input': content}) response.raise_for_status() data = response.json() - return data['data'][0]['embedding'], data['usage']['total_tokens'], anchor \ No newline at end of file + return data['data'][0]['embedding'], data['usage']['total_tokens'], anchor + + def get_word_data(self, extension: str, word_data: bytes) -> str: + """ + Extracting text from word-file + :param extension: extension of uploaded word file + :param word_file: uploaded word file + :return: word-file content + """ + try: + if extension == 'docx': + text = docx2txt.process(BytesIO(word_data)) + elif extension == 'doc': + process = subprocess.Popen( + ['antiword', '-w', '0', '-'], + stdin=subprocess.PIPE, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + ) + text, _ = process.communicate(input=word_data) + text = text.decode('utf-8') + else: + text = '' + except Exception: + text = 'Файл поврежден или не может быть прочитан.' + if text.strip(): + return f'Это текст, извлечённый из загруженного WORD-файла:\n{text}' + else: + return 'Файл пуст или содержит изображения, из которых невозможно извлечь текст.' + @@ -5,6 +5,14 @@ from decimal import Decimal from io import BytesIO from typing import Any +from replicate.exceptions import ModelError +from ml_model.exceptions import ( + RequestBlocked, + GenerationException, + ExceededContextLengthError, + ImageAnalysisError, +) + import filetype import requests from django.core.files import File @@ -63,7 +71,16 @@ class Reve(SimpleService): input_message.file.close() callback_data.update({'image': image}) type = 'edit-fast' - image = replicate_run(f'reve/{type}', callback_data) + try: + image = replicate_run(f'reve/{type}', callback_data) + except ModelError as exc: + if any(error in str(exc) for error in ('E005', 'E006', 'sexual')): + raise RequestBlocked + if 'INPUT_ANALYSIS_FAILURE' in str(exc): + raise ImageAnalysisError + if 'PROMPT_TOO_LONG' in str(exc): + raise ExceededContextLengthError + raise GenerationException from exc process_time = timedelta(seconds=(time.time() - start_time)) self.handle_invoice(input_message.content_object.model, type=type) msgs = self.save_results(input_message.content, image, process_time, save) @@ -10,6 +10,7 @@ import requests from django.core.files import File from messages.models import Message +from ml_model.exceptions import FileExtensionNotSupported, CorruptedFileError from ml_model.services.base import SimpleService from ml_model.tasks import replicate_run @@ -51,9 +52,13 @@ class Wan_Lite(SimpleService): } ) if input_message.file: - kind = filetype.guess(input_message.file.read(20)) + kind = filetype.guess(input_message.file.read(50)) + if not kind: + raise CorruptedFileError mime = kind.mime if kind else 'application/octet-stream' input_message.file.seek(0) + if kind.extension.upper() not in (available_extensions := ('JPG', 'JPEG', 'PNG', 'WEBP')): + raise FileExtensionNotSupported(available_extensions) image = f'data:{mime};base64,{base64.b64encode(input_message.file.read()).decode("utf-8")}' input_message.file.close() callback_data.update({'image': image}) @@ -1,3 +1,5 @@ +from typing import Iterable + from django.utils.translation import gettext as _ # накинуть перевод через gettext_lazy @@ -44,15 +46,28 @@ class ModelTimeoutError(Exception): class FileExtensionNotSupported(Exception): - def __init__(self, extensions: list[str]) -> None: + def __init__(self, extensions: Iterable[str]) -> None: self.extensions = extensions def __str__(self) -> str: return _( - f'The attached file format is not supported. Available formats: %(available_extensions)s.' + 'The attached file format is not supported. Available formats: %(available_extensions)s.' ) % {'available_extensions': ', '.join(self.extensions)} +class CorruptedFileError(Exception): + def __str__(self) -> str: + return _('The file may be corrupted. Please try another one.') + + +class FileTooLargeError(Exception): + def __init__(self, max_mb_size: int) -> None: + self.max_mb_size = max_mb_size + + def __str__(self) -> str: + return _('The file size cannot exceed %(max_mb_size)d MB') % {'max_mb_size': self.max_mb_size} + + class ExceededContextLengthError(Exception): def __str__(self) -> str: return _('The length of the context has been exceeded.') @@ -86,6 +101,11 @@ class ImageContentNotFound(Exception): return _('No image content found in response. Try a different request') +class ImageAnalysisError(Exception): + def __str__(self): + return _('Image analysis error. Please try another image.') + + class InvalidStyleCombinationError(Exception): def __str__(self) -> str: return _('Use style type AUTO or GENERAL when a style preset is selected') @@ -111,3 +131,8 @@ class PromptLengthExceeded(Exception): return _('Prompt is too long. Maximum length is %(max_length)s characters.') % { 'max_length': self.max_length } + + +class ServiceHighDemandError(Exception): + def __str__(self) -> str: + return _('Service is currently unavailable due to high demand. Please try again later') @@ -123,7 +123,9 @@ def openrouter_run(version: str, messages: list, callback_data: dict, model_name ) reasoning = re.sub(r'Вывод:|Основная мысль:|Рассуждение:|\*\*', '', reasoning) answer = reasoning - if 'google/gemini' in data['model']: + if any(m in data['model'] for m in ('google/gemini', 'x-ai/grok-4.1-fast')) or re.match( + r'^qwen/qwen3\.5-.*$', data['model'] + ): answer = content elif reasoning and content: # TODO: переделать рендеринг сообщения на Jinja 2 @@ -63,28 +63,34 @@ def count_openrouter_tokens(model_name: str, messages: List[Dict[str, Any]], out def create_redis_search_index() -> None: - ''' + """ A method for creating an index for storing a chunk's data (content, vectors, etc.) - ''' + """ redis_client = redis.Redis(host=settings.REDIS_HOST, port=settings.REDIS_PORT, db=0) - try: - redis_client.ft('ml_model-index').info() - except: - message_uid = TagField('message_uid') - chunk_id = TextField('chunk_id') - section_text = TextField('section_text') - section_embeddings = VectorField( - 'section_embeddings', - 'FLAT', - { - 'TYPE': 'FLOAT32', - 'DIM': 3072, - 'DISTANCE_METRIC': 'COSINE', - 'INITIAL_CAP': 10_000 - } - ) - fields = [message_uid, chunk_id, section_text, section_embeddings] - redis_client.ft('ml_model-index').create_index( - fields=fields, - definition=IndexDefinition(prefix=['ml_model:messages:'], index_type=IndexType.HASH) - ) + index_configs = ( + ('ml_model-index', 3072), + ('ml_model-index-1536', 1536), + ) + + for index_name, dim in index_configs: + try: + redis_client.ft(index_name).info() + except Exception: + message_uid = TagField('message_uid') + chunk_id = TextField('chunk_id') + section_text = TextField('section_text') + section_embeddings = VectorField( + 'section_embeddings', + 'FLAT', + { + 'TYPE': 'FLOAT32', + 'DIM': dim, + 'DISTANCE_METRIC': 'COSINE', + 'INITIAL_CAP': 10_000, + }, + ) + fields = [message_uid, chunk_id, section_text, section_embeddings] + redis_client.ft(index_name).create_index( + fields=fields, + definition=IndexDefinition(prefix=['ml_model:messages:'], index_type=IndexType.HASH), + ) @@ -27,6 +27,9 @@ from ml_model.exceptions import ( TemplateUnknownException, RequestBlocked, PromptLengthExceeded, + CorruptedFileError, + FileTooLargeError, + ImageAnalysisError, ) from ml_model.services.base import SimpleService from payments.exceptions.insufficient_balance import InsufficientBalance @@ -169,6 +172,9 @@ class MessagesAPIView(APIView): ExceededContextLengthError, RequestBlocked, PromptLengthExceeded, + CorruptedFileError, + FileTooLargeError, + ImageAnalysisError, ) as exc: return Response({'detail': f'{exc}'}, status=HTTP_400_BAD_REQUEST) except TemplateNotFound as exc: @@ -18,7 +18,12 @@ from ml_model.exceptions import ( InvalidParameterError, InvalidStyleCombinationError, PromptLengthExceeded, + ExceededContextLengthError, FileExtensionNotSupported, + ServiceHighDemandError, + CorruptedFileError, + FileTooLargeError, + ImageAnalysisError, ) from ml_model.models import NeuronModel from ml_model.services.base import SimpleService @@ -159,6 +164,17 @@ class MediaAPIView(APIView): logger.exception(exc) input_message.is_sent = False input_message.save() + if any( + phrase in str(exc) for phrase in ('Insufficient credit', 'Request was throttled') + ): + return Response( + { + 'detail': _( + 'Temporary issues with the service, we are already working on a solution.' + ) + }, + status=HTTP_402_PAYMENT_REQUIRED, + ) if isinstance(exc, InsufficientBalance): return Response({'detail': f'{exc}'}, status=HTTP_402_PAYMENT_REQUIRED) if isinstance( @@ -171,7 +187,12 @@ class MediaAPIView(APIView): InvalidStyleCombinationError, InvalidParameterError, PromptLengthExceeded, + ExceededContextLengthError, FileExtensionNotSupported, + ServiceHighDemandError, + CorruptedFileError, + FileTooLargeError, + ImageAnalysisError, ), ): return Response({'detail': f'{exc}'}, status=HTTP_400_BAD_REQUEST)