@@ -346,13 +346,6 @@ YANDEX_CLOUD_ID = env.str('YANDEX_CLOUD_ID', 'defaultapikey') OPENAI_PROXY_HOST = env.str('OPENAI_PROXY_HOST', 'neuron-proxy:8080') UPSCALE_MULTIPLIER_HOST = env.str('UPSCALE_MULTIPLIER_HOST', 'packet:8080') - -MAX_UPLOAD_SIZE_PER_MODEL = { - 'raifgpt': 50, - 'default': 8, -} - - # Payments YOOKASSA_ACCOUNT_ID = env.str('YOOKASSA_ACCOUNT_ID', default='defaultapikey') YOOKASSA_SECRET_KEY = env.str('YOOKASSA_SECRET_KEY', default='defaultapikey') @@ -1323,3 +1323,6 @@ msgstr "" msgid "You cannot change the password of an unconfirmed e-mail user." msgstr "Вы не можете изменить пароль неподтвержденного по e-mail пользователя." + +msgid "Unknown file format" +msgstr "Неизвестный формат файла" @@ -1,4 +1,10 @@ +from typing import Dict, Any + +from backend import settings + +from django.utils.translation import gettext_lazy as _ from rest_framework import serializers +from rest_framework.serializers import ValidationError from messages.models import Message @@ -31,3 +37,12 @@ class MessageSerializer(serializers.ModelSerializer): 'is_sent', 'info', ] + + def validate(self, data: Dict[str, Any]) -> Dict[str, Any]: + file = data.get('file') + max_mb_size = 52 + if file and file.size > (max_mb_size << 10 << 10): + raise ValidationError( + _('The file size cannot exceed %(max_mb_size)d MB') % {'max_mb_size': max_mb_size} + ) + return data @@ -1,14 +1,15 @@ import base64 import json import logging +import httpx +import filetype import re import uuid + from abc import ABC, abstractmethod from io import BytesIO, StringIO from typing import Any, Iterable, Literal -import filetype -import httpx from django.conf import settings from django.core.exceptions import ValidationError from django.utils.translation import gettext_lazy as _ @@ -42,6 +43,7 @@ class OpenAICompatibleRunner(BaseRunner, ABC): @classmethod def generate(cls, content=None, file=None, parameters={}, history=[], scrape_results=[]): proxies = Proxy.objects.all() + system_prompt = parameters.pop('system_prompt') payload = { 'messages': [ *[ @@ -56,6 +58,9 @@ class OpenAICompatibleRunner(BaseRunner, ABC): 'stream': True, **parameters, } + + payload['messages'].insert(0, {'role': 'system', 'content': system_prompt}) + for result in scrape_results: if isinstance(result, StringIO): payload['messages'].append( @@ -3,18 +3,23 @@ import logging import re import subprocess import time +import zipfile +import fitz +import httpx from datetime import timedelta from decimal import Decimal from io import BytesIO, StringIO from math import ceil -from typing import Any, Callable, Iterable, Literal +from typing import Any, Callable, Literal, Tuple, Union, Optional, Iterable from uuid import UUID import docx2txt +import openpyxl import filetype from django.core.cache import cache from django.core.files.base import ContentFile -from django.db.models import Prefetch +from django.conf import settings +from django.db.models import Prefetch, QuerySet from django.utils.translation import gettext_lazy as _ from PIL import Image as ImageModule from PyPDF2 import PdfReader @@ -24,7 +29,7 @@ from messages.models import Message from ml_model.exceptions import ( InferenceDisabled, PaymentRuleNotImplemented, - ScraperDoesNotExists, + ScraperDoesNotExists, UnknownFileException, ) from ml_model.models import ( Deployment, @@ -45,6 +50,9 @@ logger = logging.getLogger(__name__) class InferenceService: def __init__(self, user: CustomUserModel): self.user = user + self.DATA_SIZE_MULTIPLIERS = { + 'mb': 1024 * 1024, + } @classmethod def get_by_id(cls, id: UUID): @@ -117,6 +125,20 @@ class InferenceService: raw_file.seek(0) file_buf = BytesIO(raw_file.read()) + if file_extension == 'zip': + file_extension = None + signatures = { + 'xlsx': 'xl/workbook.xml', + 'docx': 'word/document.xml' + } + with zipfile.ZipFile(file_buf, 'r') as zip_file: + namelist = zip_file.namelist() + for format_name, required_file in signatures.items(): + if required_file in namelist: + file_extension = format_name + break + if not file_extension: + raise UnknownFileException if file_extension in ('png', 'jpg', 'jpeg'): file = BytesIO() normalized_image = ImageModule.open(file_buf) @@ -124,10 +146,12 @@ class InferenceService: file.seek(0) elif file_extension in ('pdf',): file = StringIO() - reader = PdfReader(file_buf) - for page in reader.pages: - file.write(page.extract_text()) - elif file_extension in ('doc', 'docx', 'zip'): + pdf_data = self.get_pdf_data(file_buf, raw_file.name) + raw_text = pdf_data[0] + image_count = pdf_data[1] + file.write(raw_text) + file.seek(0) + elif file_extension in ('doc', 'docx'): extractors: dict[Literal['doc', 'docx'], Callable[[], str]] = { 'doc': lambda: subprocess.Popen( ['antiword', '-w', '0', '-'], @@ -138,13 +162,20 @@ class InferenceService: .communicate(file_buf.getvalue()) .decode(), 'docx': lambda: docx2txt.process(file_buf), - 'zip': lambda: docx2txt.process(file_buf), } file = StringIO() file.write('Remember this Document included in request:') file.write('[DOCUMENT-START]\n') file.write(extractors[file_extension]()) file.write('\n[DOCUMENT-END]') + elif file_extension in ('xlsx',): + file = StringIO() + xlsx_file = openpyxl.load_workbook(file_buf) + for sheet_name in xlsx_file.sheetnames: + sheet = xlsx_file[sheet_name] + for row in sheet.iter_rows(values_only=True): + file.write(f'Данные ряда: {row}') + except Exception: file = None logger.info(f'File: {file}') @@ -263,8 +294,16 @@ class InferenceService: ) if scrape_results: calculated_price += ( - sum([TokenizerTool.token_count(result.getvalue()) for result in scrape_results]) + sum([TokenizerTool.token_count(result.getvalue()) for result in + scrape_results]) ) * payment_rule.cost + if (system_prompt := parameters.get('system_prompt')): + calculated_price += ( + sum([TokenizerTool.token_count(prompt) for prompt in system_prompt]) + * payment_rule.cost + ) + if image_count: + calculated_price += image_count * Decimal('0.13') elif ( payment_rule.strategy == PaymentRule.StrategyChoices.PER_TEXT_TOKEN and payment_rule.interaction_type == PaymentRule.InteractionTypeChoices.OUTPUT @@ -375,3 +414,191 @@ class InferenceService: logger.exception(exc) finally: cache.delete(cache_key) + + def get_pdf_data(self, input_data: Union[Message, BytesIO], filename: Optional[str] = None) -> str: + try: + if isinstance(input_data, BytesIO): + file_stream = input_data + actual_filename = filename or getattr(input_data, 'name', '') + else: + file_stream = input_data.file + actual_filename = filename or getattr(input_data.file, 'name', '') + processor = PDFProcessor() + result, image_count = processor.process( + file_stream=file_stream, + filename=actual_filename, + ) + return result, image_count + except Exception as exc: + logger.exception(exc) + return f"Ошибка обработки PDF: {str(exc)}" + finally: + if isinstance(input_data, BytesIO): + input_data.seek(0) + elif hasattr(input_data, 'file'): + input_data.file.seek(0) + + +class PDFProcessor: + MAX_UPLOAD_SIZE_PER_MODEL: dict = { + 'mb': 1024 * 1024, + } + MAX_BATCH_SIZE = 3.9 * MAX_UPLOAD_SIZE_PER_MODEL['mb'] + + def __init__(self): + self.image_count = 0 + + def process(self, file_stream: BytesIO, filename: str = "") -> Tuple[str, int]: + try: + file_stream.seek(0) + pdf_data = file_stream.read() + doc = fitz.open(stream=pdf_data, filetype="pdf") + has_images = any(page.get_images() for page in doc) + if has_images: + text = self._process_ocr(doc) + return text, self.image_count + raw_text = [] + for page in doc: + content = page.get_text("text") + if content: + raw_text.append(content) + return "\n".join(raw_text), 0 + finally: + doc.close() + fitz.TOOLS.store_shrink(100) + + def _process_ocr(self, doc) -> str: + raw_texts, pages_with_image = self._extract_text_and_images(doc) + ocr_texts = self._process_images_with_yandex_vision(doc, pages_with_image) + final_text = self._combine_texts(raw_texts, ocr_texts) + return final_text + + def _extract_text_and_images(self, doc) -> Tuple[dict, list]: + raw_texts = {} + pages_with_image = [] + for page_num, page in enumerate(doc): + text = page.get_text("text") + if text: + raw_texts[page_num] = text + if page.get_images(): + pages_with_image.append(page_num) + return raw_texts, pages_with_image + + def _process_images_with_yandex_vision(self, doc, page_nums) -> dict: + batch_images, page_index_map = self._prepare_image_batches(doc, page_nums) + self.image_count = len(batch_images) + if not batch_images: + return {} + return self._send_to_yandex_vision(batch_images, page_index_map) + + def _prepare_image_batches(self, doc, page_nums) -> Tuple[list, list]: + batch_images = [] + page_index_map = [] + for page_num in page_nums: + try: + page = doc.load_page(page_num) + pix = page.get_pixmap(dpi=150, alpha=False) + img = ImageModule.frombytes("RGB", [pix.width, pix.height], pix.samples) + buffer = BytesIO() + img.save(buffer, format="JPEG", quality=60, optimize=True) + buffer.seek(0) + if buffer.getbuffer().nbytes < self.__class__.MAX_BATCH_SIZE: + batch_images.append(buffer) + page_index_map.append(page_num) + except Exception as e: + logger.error(f"Error processing page {page_num}: {e}") + continue + return batch_images, page_index_map + + def _send_to_yandex_vision(self, batch_images, page_index_map) -> dict: + headers = { + "Authorization": f"Api-Key {settings.YANDEX_CLOUD_API_KEY}", + "Content-Type": "application/json" + } + ocr_results = {} + batches = self._create_batches(batch_images, page_index_map) + for batch, pages in batches: + body = { + "folderId": settings.YANDEX_CLOUD_ID, + "analyze_specs": [{ + "content": base64.b64encode(buf.getvalue()).decode(), + "features": [{ + "type": "TEXT_DETECTION", + "text_detection_config": {"language_codes": ["*"]} + }] + } for buf in batch] + } + try: + resp = httpx.post( + "https://vision.api.cloud.yandex.net/vision/v1/batchAnalyze", + headers=headers, json=body, timeout=60 + ) + if resp.status_code >= 400: + self.image_count = 0 + return {} + response = resp.json() if resp.status_code == 200 else None + if response: + self._parse_vision_response(response, pages, ocr_results) + except Exception as e: + logger.error(f"Yandex Vision API error: {e}") + return {} + return ocr_results + + def _create_batches(self, batch_images, page_index_map) -> list: + batches = [] + current_batch = [] + current_pages = [] + current_size = 0 + for i, buffer in enumerate(batch_images): + size = buffer.getbuffer().nbytes + if current_size + size > self.__class__.MAX_BATCH_SIZE and current_batch: + batches.append((current_batch, current_pages)) + current_batch, current_pages, current_size = [], [], 0 + current_batch.append(buffer) + current_pages.append(page_index_map[i]) + current_size += size + if current_batch: + batches.append((current_batch, current_pages)) + return batches + + def _make_yandex_vision_request(self, batch, headers): + body = { + "folderId": settings.YANDEX_CLOUD_ID, + "analyze_specs": [{ + "content": base64.b64encode(buf.getvalue()).decode(), + "features": [{ + "type": "TEXT_DETECTION", + "text_detection_config": {"language_codes": ["*"]} + }] + } for buf in batch] + } + try: + resp = httpx.post( + "https://vision.api.cloud.yandex.net/vision/v1/batchAnalyze", + headers=headers, json=body, timeout=60 + ) + return resp.json() if resp.status_code == 200 else None + except Exception as e: + logger.error(f"Yandex Vision API error: {e}") + return None + + def _parse_vision_response(self, response, pages, ocr_results): + for i, spec_result in enumerate(response.get("results", [])): + page_text = [] + for res in spec_result.get("results", []): + for page in res.get("textDetection", {}).get("pages", []): + for block in page.get('blocks', []): + for line in block.get('lines', []): + line_text = " ".join( + word.get('text', '') for word in line.get('words', []) + ) + if line_text: + page_text.append(line_text) + ocr_results[pages[i]] = "\n".join(page_text) + + def _combine_texts(self, raw_texts, ocr_texts) -> str: + all_pages = sorted(set(raw_texts) | set(ocr_texts)) + return "\n\n".join( + f"{raw_texts.get(pn, '')}\n{ocr_texts.get(pn, '')}".strip() + for pn in all_pages + ).strip() or "Не удалось распознать текст" @@ -39,3 +39,8 @@ class FileExtensionNotSupported(Exception): return _( 'The attached file format is not supported. Available formats: %(available_extensions)s.' ) % {'available_extensions': ', '.join(self.extensions)} + + +class UnknownFileException(Exception): + def __str__(self): + return _('Unknown file format') \ No newline at end of file