@@ -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,15 @@ 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', + ): 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 +127,28 @@ 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, + index_name=index_name, ) ] drop_redis_vectors.delay(message_uid) @@ -1,30 +1,41 @@ 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 messages.models import Message +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.media.models import Image 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 +54,63 @@ 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]) + 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}' + ) + else: + 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}}, + ] 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 +118,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 @@ -123,7 +123,7 @@ 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')): 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), + )