@@ -0,0 +1,130 @@ +import httpx +import numpy as np +import redis + +from concurrent.futures import ThreadPoolExecutor, as_completed +from typing import Tuple, List + +from django.conf import settings +from langchain_text_splitters import RecursiveCharacterTextSplitter +from pydantic.v1 import UUID4 +from redis.commands.search.document import Document +from redis.commands.search.query import Query + +from ml_model.tasks import drop_redis_vectors + + +class EmbeddingService: + @classmethod + def split_text_to_chunks(cls, raw_text: str, chunk_size: int = 4000, overlap: int = 200) -> list[str]: + text_splitter = RecursiveCharacterTextSplitter( + chunk_size=chunk_size, + chunk_overlap=overlap, + length_function=len, + separators=['\n\n', '\n', '.', ' ', ''], + ) + return text_splitter.split_text(raw_text) + + @classmethod + def process_chunk( + cls, + client: httpx.Client, + chunk: str, + redis_client: redis.Redis, + message_uid: str, + chunk_id: int, + ) -> int: + embedding, e_total_tokens = cls._get_embedding(client=client, content=chunk) + cls._save_embeddings( + redis_client=redis_client, + message_uid=message_uid, + chunk_id=chunk_id, + text=chunk, + embeddings=embedding, + ) + 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}) + response.raise_for_status() + data = response.json() + return data['data'][0]['embedding'], data['usage']['total_tokens'] + + @classmethod + def _save_embeddings( + cls, redis_client: redis.Redis, message_uid: str, chunk_id: int, text: str, embeddings: List[float] + ) -> None: + 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, + }, + ) + + @classmethod + def search_via_embeddings( + cls, + redis_client: redis.Redis, + message_uid: str, + user_query_embeddings: List[float], + top_k: int = 10, + ) -> List[Document]: + 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 + + @classmethod + def make_embeddings_prompt(cls, document_name: str, section_texts: List[str], question: str) -> str: + return f"""Ты — аналитик данных. Отвечай только на основе предоставленного контекста. + Название файла: {document_name} + Фрагменты: + {'\n'.join(section_texts)} + Вопрос: {question} + """ + + @classmethod + def get_large_file_data(cls, msg_uid: UUID4, chunks, proxy, user_content): + redis_client = redis.Redis(host=settings.REDIS_HOST, port=settings.REDIS_PORT, db=0) + embedding_tokens = 0 + message_uid = str(msg_uid).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( + cls.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 = cls._get_embedding(client=client, content=user_content) + 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 + ) + ] + drop_redis_vectors.delay(message_uid) + redis_client.close() + return embedding_tokens, result \ No newline at end of file @@ -0,0 +1,85 @@ +import re +import subprocess +import zipfile +import docx2txt +import fitz +import openpyxl + +from io import BytesIO + + +class FileProcessingService: + @classmethod + def get_file_extension(cls, 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 + + @classmethod + def get_file_data(cls, file_extension: str, file_bytes: bytes) -> str: + is_word = file_extension in ('doc', 'docx') + method_name = 'word' if is_word else file_extension + operation = getattr(cls, 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 text + + @classmethod + def get_pdf_data(cls, pdf_data: bytes) -> str: + 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()}' + + @classmethod + def get_xlsx_data(cls, xlsx_data: bytes) -> str: + 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}' + + @classmethod + def get_word_data(cls, extension: str, word_data: bytes) -> str: + 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 'Файл пуст или содержит изображения, из которых невозможно извлечь текст.' + @@ -10,8 +10,12 @@ 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 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 @@ -32,7 +36,7 @@ class Claude(SimpleService): 'claude-3.5-haiku': { 'input': Decimal('1200'), 'output': Decimal('1200'), - }, # 1M tokens + }, # 1M tokens 'claude-sonnet-4.5': { 'input': Decimal('900'), 'output': Decimal('4500'), @@ -43,8 +47,10 @@ class Claude(SimpleService): }, # 1M tokens } + TOOLS_TOKEN_COSTS = {'text-embedding-3-large': {'output': Decimal('0.000065')}} + def calculate_price( - self, version: str, input_tokens: int, output_tokens: int, image: FieldFile + self, version: str, input_tokens: int, output_tokens: int, image: FieldFile, embedding_tokens: int ) -> Decimal: price_map = self.TOKENS_COST[version.split('/')[1]] price = ( @@ -52,6 +58,8 @@ class Claude(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 return price.quantize(Decimal('0.1'), rounding='ROUND_UP') def save_results(self, content: Iterator[Any], t: timedelta, save: bool = True) -> list[Message]: @@ -74,22 +82,55 @@ class Claude(SimpleService): messages = [ {'role': 'system', 'content': system_prompt}, *self.get_chat_history(), - {'role': 'user', 'content': input_message.content} + {'role': 'user', 'content': input_message.content}, ] - image = input_message.file - if image: - kind = filetype.guess(image.read(20)) - mime = kind.mime if kind else 'application/octet-stream' - normalized_image = Image.open(image) - 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}}, - ] + file = input_message.file + image = None + embedding_tokens = 0 + if file: + file_service = FileProcessingService + try: + 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 + ) + 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(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}}, + ] + image = file + except Exception: + raise FileExtensionNotSupported(['PDF', 'DOC', 'DOCX', 'XLSX', 'JPG', 'JPEG', 'PNG']) result = openrouter_run(version, messages, callback_data, 'Claude') process_time = timedelta(seconds=(time.time() - start_time)) self.handle_invoice( @@ -98,17 +139,20 @@ class Claude(SimpleService): input_tokens=result[1], output_tokens=result[2], image=image, + 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]]: + 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] + ).order_by('-created_at')[1 : message_limit + 1] ) ) elif isinstance(self.store, APIStore): @@ -9,8 +9,11 @@ from django.db.models.fields.files import FieldFile 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.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 @@ -55,8 +58,10 @@ class Gemini(SimpleService): }, } + TOOLS_TOKEN_COSTS = {'text-embedding-3-large': {'output': Decimal('0.000065')}} + def calculate_price( - self, version: str, input_tokens: int, output_tokens: int, image: FieldFile + 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: @@ -71,6 +76,9 @@ class Gemini(SimpleService): ) if image: price += price_map['input_imgs'] / 1_000 + if embedding_tokens > 0: + print(embedding_tokens) + price += self.TOOLS_TOKEN_COSTS['text-embedding-3-large']['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]: @@ -94,20 +102,50 @@ class Gemini(SimpleService): } messages = self.get_chat_history() messages.append({'role': 'user', 'content': input_message.content}) - image = input_message.file - if image: - kind = filetype.guess(image.read(20)) - mime = kind.mime if kind else 'application/octet-stream' - normalized_image = Image.open(image) - 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}}, - ] + file = input_message.file + image = None + embedding_tokens = 0 + if 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 + ) + 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(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}}, + ] + image = file start_time = time.time() result = openrouter_run(version, messages, callback_data, 'Gemini') process_time = timedelta(seconds=(time.time() - start_time)) @@ -117,6 +155,7 @@ class Gemini(SimpleService): input_tokens=result[1], output_tokens=result[2], image=image, + embedding_tokens=embedding_tokens, ) msgs = self.save_results(result[0], process_time) return msgs