@@ -3,16 +3,18 @@ import time from datetime import timedelta from decimal import Decimal from io import BytesIO +from pathlib import Path from typing import Any, Iterator import filetype from PIL import Image from messages.models import Message -from ml_model.exceptions import FileExtensionNotSupported, FileUploadUnsupported, ModelVersionNotAvailable +from ml_model.adapters.openrouter import OpenrouterAdapter +from ml_model.exceptions import CorruptedFileError, FileExtensionNotSupported, ModelVersionNotAvailable from ml_model.services.EmbeddingService import EmbeddingService from ml_model.services.FileService import FileProcessingService -from ml_model.services.base import SimpleService +from ml_model.services.base import StreamSimpleService from ml_model.tasks import openrouter_run from payments.exceptions.insufficient_balance import InsufficientBalance from payments.selectors.payment_plan_selector import PaymentPlanSelector @@ -21,8 +23,7 @@ from tools.chats.models import Chat from tools.copywrite.models import Copywrite from tools.public_api.models import APIStore - -class Claude(SimpleService): +class Claude(StreamSimpleService): """ Claude Service contains abstract method make, which makes a generation @@ -71,10 +72,12 @@ class Claude(SimpleService): TOOLS_TOKEN_COSTS = {'text-embedding-3-small': {'output': Decimal('0.00001')}} + SUPPORTED_EXTENSIONS = ['PDF', 'DOC', 'DOCX', 'XLSX', 'JPG', 'JPEG', 'PNG', 'WEBP'] + def calculate_price( self, version: str, input_tokens: int, output_tokens: int, embedding_tokens: int ) -> Decimal: - price_map = self.TOKENS_COST[version.split('/')[1]] + price_map = self.TOKENS_COST[version] price = ( input_tokens * price_map['input'] / 1_000_000 + output_tokens * price_map['output'] / 1_000_000 ) @@ -83,7 +86,7 @@ class Claude(SimpleService): price += Decimal('2') return price.quantize(Decimal('0.1'), rounding='ROUND_UP') - def save_results(self, content: Iterator[Any], t: timedelta, save: bool = True) -> list[Message]: + def save_results(self, content: str, t: timedelta, save: bool = True) -> list[Message]: msgs = [ Message( content=content, @@ -97,102 +100,61 @@ class Claude(SimpleService): def make(self, input_message: Message, save: bool = True) -> list[Message]: start_time = time.time() - version_slug = input_message.info.pop('version', None) + version_slug = input_message.info.get('version') if version_slug is None or version_slug not in self.TOKENS_COST: raise ModelVersionNotAvailable(version_slug, self.TOKENS_COST) - version = f'anthropic/{version_slug}' - system_prompt = input_message.info.pop('system_prompt', '') - callback_data = {'provider': {'order': ['anthropic']}, **input_message.info, 'tools': []} - messages = [ - {'role': 'system', 'content': system_prompt}, - *self.get_chat_history(), - {'role': 'user', 'content': input_message.content}, - ] - messages.insert(0, {'role': 'system', 'content': self.HISTORY_CONTEXT_SYSTEM_PROMPT}) - if version_slug == 'claude-fable-5': - current_user_balance = PaymentPlanSelector(self.store.user).get_current_balance() - if current_user_balance < (cost := Decimal('100')) and input_message.file: - raise InsufficientBalance(current_user_balance, cost) - messages.insert(0, {'role': 'system', 'content': self.FABLE_SYSTEM_PROMPT}) - if reasoning := input_message.info.get('reasoning'): - reasoning_data = { - 'Средний': 'low', - 'Высокий': 'medium', - } - callback_data['reasoning'] = {'effort': reasoning_data[reasoning]} - callback_data['tools'].append( - { - 'type': 'openrouter:web_search', - 'parameters': { - 'engine': 'parallel', - 'max_results': 1, - 'max_total_results': 3, - 'search_context_size': 'low', - }, - } - ) - file = input_message.file - embedding_tokens = 0 - if file: - file_service = FileProcessingService - try: - file_bytes = input_message.file.read() - kind = filetype.guess(file_bytes[:550]) - raw_file_extension = kind.extension - file_extension = file_service.get_file_extension(raw_file_extension, file_bytes) - if file_extension in ('pdf', 'doc', 'docx', 'xlsx'): - text = file_service.get_file_data(file_extension, file_bytes) - chunks = EmbeddingService.split_text_to_chunks(text) - if len(text) > 20_000: - for proxy in Proxy.objects.all(): - document_name = ( - chunks[0].partition(f':{chr(10)}')[2].split(f'{chr(10)}')[0][:100] - ) - embedding_tokens, file_data = EmbeddingService.get_large_file_data( - self.store.messages.first().pk, - chunks, - proxy, - input_message.content, - model='text-embedding-3-small', - index_name='ml_model-index-1536', - ) - messages[-1]['content'] = EmbeddingService.make_embeddings_prompt( - document_name=document_name, - section_texts=file_data, - question=input_message.content, - ) - else: - messages[-1]['content'] = ( - f'Используй системный промпт. Содержание файла: ' - f'{chunks}. Вопрос: {input_message.content}' - ) - else: - kind = filetype.guess(file_bytes[:20]) - mime = kind.mime if kind else 'application/octet-stream' - format = 'jpeg' if kind.extension == 'jpg' else kind.extension - with Image.open(file) as normalized_image: - with BytesIO() as buf: - normalized_image.save(buf, format=format) - image_url = ( - f'data:{mime};base64,{base64.b64encode(buf.getvalue()).decode("utf-8")}' - ) - messages[-1]['content'] = [ - {'type': 'text', 'text': input_message.content}, - {'type': 'image_url', 'image_url': {'url': image_url}}, - ] - except Exception: - raise FileExtensionNotSupported(['PDF', 'DOC', 'DOCX', 'XLSX', 'JPG', 'JPEG', 'PNG']) - result = openrouter_run(f'{version}', messages, callback_data, 'Claude') + model_slug = f'anthropic/{version_slug}' + callback_data = self._build_callback_data(input_message, version_slug) + messages, embedding_tokens = self._prepare_messages(input_message, version_slug) + + result = openrouter_run(model_slug, messages, callback_data, 'Claude') + process_time = timedelta(seconds=(time.time() - start_time)) self.handle_invoice( input_message.content_object.model, - version=version, + version=version_slug, input_tokens=result[1], output_tokens=result[2], embedding_tokens=embedding_tokens, ) - msgs = self.save_results(result[0], process_time) - return msgs + return self.save_results(result[0], process_time, save) + + def make_stream(self, input_message: Message, save: bool = True) -> Iterator[str]: + start_time = time.time() + version_slug = input_message.info.get('version') + if version_slug is None or version_slug not in self.TOKENS_COST: + raise ModelVersionNotAvailable(version_slug, self.TOKENS_COST) + model_slug = f'anthropic/{version_slug}' + callback_data = self._build_callback_data(input_message, version_slug) + messages, embedding_tokens = self._prepare_messages(input_message, version_slug) + content_parts: list[str] = [] + input_tokens = output_tokens = 0 + result = '' + + try: + stream = OpenrouterAdapter.run_streaming_api(model_slug, messages, callback_data, 'Claude') + try: + while True: + chunk = next(stream) + if chunk: + content_parts.append(chunk) + yield chunk + except StopIteration as exc: + input_tokens, output_tokens = exc.value + finally: + if content_parts: + result = ''.join(content_parts) + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice( + input_message.content_object.model, + version=version_slug, + input_tokens=input_tokens, + output_tokens=output_tokens, + embedding_tokens=embedding_tokens, + ) + self.save_results(result, process_time, save) + if result: + return result def get_chat_history( self, message_limit: int = 10, max_character_limit: int = 1500 @@ -217,6 +179,8 @@ class Claude(SimpleService): ).order_by('-created_at')[:message_limit] ) ) + else: + air_messages = [] memory = [] for msg in air_messages: content = msg.content or '' @@ -229,3 +193,95 @@ class Claude(SimpleService): character_length -= len(memory.pop(0)['content']) return memory + + def _build_callback_data(self, input_message: Message, version_slug: str) -> dict[str, Any]: + callback_data = {'provider': {'order': ['anthropic']}, **input_message.info, 'tools': []} + if version_slug == 'claude-fable-5': + if reasoning := input_message.info.get('reasoning'): + reasoning_data = { + 'Средний': 'low', + 'Высокий': 'medium', + } + callback_data['reasoning'] = {'effort': reasoning_data[reasoning]} + callback_data['tools'].append( + { + 'type': 'openrouter:web_search', + 'parameters': { + 'engine': 'parallel', + 'max_results': 1, + 'max_total_results': 3, + 'search_context_size': 'low', + }, + } + ) + return callback_data + + def _prepare_messages( + self, input_message: Message, version_slug: str + ) -> tuple[list[dict[str, str | list]], int]: + if version_slug == 'claude-fable-5' and input_message.file: + current_user_balance = PaymentPlanSelector(self.store.user).get_current_balance() + if current_user_balance < (cost := Decimal('100')): + raise InsufficientBalance(current_user_balance, cost) + + system_prompt = input_message.info.get('system_prompt', '') + messages = [ + {'role': 'system', 'content': self.HISTORY_CONTEXT_SYSTEM_PROMPT}, + {'role': 'system', 'content': system_prompt}, + *self.get_chat_history(), + {'role': 'user', 'content': input_message.content}, + ] + if version_slug == 'claude-fable-5': + messages.insert(1, {'role': 'system', 'content': self.FABLE_SYSTEM_PROMPT}) + + embedding_tokens = 0 + if input_message.file: + file_service = FileProcessingService + file_bytes = input_message.file.read() + kind = filetype.guess(file_bytes[:550]) + if not kind: + if Path(input_message.file.name).suffix[1:].upper() not in self.SUPPORTED_EXTENSIONS: + raise FileExtensionNotSupported(self.SUPPORTED_EXTENSIONS) + raise CorruptedFileError + raw_file_extension = kind.extension + file_extension = file_service.get_file_extension(raw_file_extension, file_bytes) + if file_extension in ('pdf', 'doc', 'docx', 'xlsx'): + text = file_service.get_file_data(file_extension, file_bytes) + chunks = EmbeddingService.split_text_to_chunks(text) + if len(text) > 20_000: + for proxy in Proxy.objects.all(): + document_name = chunks[0].partition(f':{chr(10)}')[2].split(f'{chr(10)}')[0][:100] + embedding_tokens, file_data = EmbeddingService.get_large_file_data( + self.store.messages.first().pk, + chunks, + proxy, + input_message.content, + model='text-embedding-3-small', + index_name='ml_model-index-1536', + ) + messages[-1]['content'] = EmbeddingService.make_embeddings_prompt( + document_name=document_name, + section_texts=file_data, + question=input_message.content, + ) + else: + messages[-1]['content'] = ( + f'Используй системный промпт. Содержание файла: ' + f'{chunks}. Вопрос: {input_message.content}' + ) + elif file_extension in ('jpg', 'jpeg', 'png', 'webp'): + kind = filetype.guess(file_bytes[:20]) + mime = kind.mime if kind else 'application/octet-stream' + format = 'jpeg' if kind.extension == 'jpg' else kind.extension + with Image.open(input_message.file) as normalized_image: + with BytesIO() as buf: + normalized_image.save(buf, format=format) + image_url = f'data:{mime};base64,{base64.b64encode(buf.getvalue()).decode("utf-8")}' + messages[-1]['content'] = [ + {'type': 'text', 'text': input_message.content}, + {'type': 'image_url', 'image_url': {'url': image_url}}, + ] + else: + raise FileExtensionNotSupported(self.SUPPORTED_EXTENSIONS) + + return messages, embedding_tokens @@ -4,16 +4,17 @@ from datetime import timedelta from decimal import Decimal from io import BytesIO from pathlib import Path +from typing import Iterator import filetype from PIL import Image from messages.models import Message -from ml_model.exceptions import CorruptedFileError, FileExtensionNotSupported +from ml_model.adapters.openrouter import OpenrouterAdapter +from ml_model.exceptions import CorruptedFileError, FileExtensionNotSupported, ModelVersionNotAvailable from ml_model.services.EmbeddingService import EmbeddingService from ml_model.services.FileService import FileProcessingService -from ml_model.exceptions import ModelVersionNotAvailable -from ml_model.services.base import SimpleService +from ml_model.services.base import StreamSimpleService from ml_model.tasks import openrouter_run from poller.models import Proxy from tools.chats.models import Chat @@ -21,7 +22,8 @@ from tools.copywrite.models import Copywrite from tools.public_api.models import APIStore -class Gemini_3_1(SimpleService): + +class Gemini_3_1(StreamSimpleService): TOKENS_COST = { 'gemini-3.1-pro-preview': { 'input': Decimal('600'), @@ -37,6 +39,8 @@ class Gemini_3_1(SimpleService): TOOLS_TOKEN_COSTS = {'text-embedding-3-small': {'output': Decimal('0.00001')}} + SUPPORTED_EXTENSIONS = ['PDF', 'DOC', 'DOCX', 'XLSX', 'JPG', 'JPEG', 'PNG', 'WEBP'] + def calculate_price(self, version: str, input_tokens: int, output_tokens: int, embedding_tokens: int) -> Decimal: price_map = self.TOKENS_COST[version] if input_tokens >= 200_000 or output_tokens >= 200_000: @@ -67,6 +71,31 @@ class Gemini_3_1(SimpleService): return msgs def make(self, input_message: Message, save: bool = True) -> list[Message]: + start_time = time.time() + version_slug = input_message.info.get('version') + if version_slug is None or version_slug not in self.TOKENS_COST: + raise ModelVersionNotAvailable(version_slug, self.TOKENS_COST) + model_slug = f'google/{version_slug}:online' + callback_data = { + 'provider': {'order': ['Google AI Studio']}, + **input_message.info, + } + messages, embedding_tokens = self._prepare_messages(input_message) + + result = openrouter_run(model_slug, messages, callback_data, 'Gemini') + + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice( + input_message.content_object.model, + version=version_slug, + input_tokens=result[1], + output_tokens=result[2], + embedding_tokens=embedding_tokens, + ) + return self.save_results(result[0], process_time, save) + + def make_stream(self, input_message: Message, save: bool = True) -> Iterator[str]: + start_time = time.time() version_slug = input_message.info.get('version') if version_slug is None or version_slug not in self.TOKENS_COST: raise ModelVersionNotAvailable(version_slug, self.TOKENS_COST) @@ -75,6 +104,76 @@ class Gemini_3_1(SimpleService): 'provider': {'order': ['Google AI Studio']}, **input_message.info, } + messages, embedding_tokens = self._prepare_messages(input_message) + content_parts: list[str] = [] + input_tokens = output_tokens = 0 + result = '' + + try: + stream = OpenrouterAdapter.run_streaming_api(model_slug, messages, callback_data, 'Gemini') + try: + while True: + chunk = next(stream) + if chunk: + content_parts.append(chunk) + yield chunk + except StopIteration as exc: + input_tokens, output_tokens = exc.value + finally: + if content_parts: + result = ''.join(content_parts) + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice( + input_message.content_object.model, + version=version_slug, + input_tokens=input_tokens, + output_tokens=output_tokens, + embedding_tokens=embedding_tokens, + ) + self.save_results(result, process_time, save) + if result: + return result + + def get_chat_history( + self, message_limit: int = 10, max_character_limit: int = 1500 + ) -> list[dict[str, str | list]]: + if isinstance(self.store, Chat): + air_messages = list( + reversed( + Message.objects.filter( + chats_chats_messages=self.store, is_deleted=False, is_sent=True + ).order_by('-created_at')[1 : message_limit + 1] + ) + ) + elif isinstance(self.store, APIStore): + air_messages = [] + elif isinstance(self.store, Copywrite): + air_messages = list( + reversed( + Message.objects.filter( + copywrite_copywrites_messages=self.store, + is_deleted=False, + is_sent=True, + ).order_by('-created_at')[:message_limit] + ) + ) + else: + air_messages = [] + memory = [] + for msg in air_messages: + content = msg.content or '' + if msg.from_model: + memory.append({'role': 'assistant', 'content': content}) + else: + memory.append({'role': 'user', 'content': content}) + character_length = sum(len(content['content']) for content in memory) + while character_length > max_character_limit: + character_length -= len(memory.pop(0)['content']) + return memory + + def _prepare_messages( + self, input_message: Message + ) -> tuple[list[dict[str, str | list]], int]: messages = self.get_chat_history() messages.insert( 0, @@ -89,13 +188,12 @@ class Gemini_3_1(SimpleService): messages.append({'role': 'user', 'content': input_message.content}) embedding_tokens = 0 if input_message.file: - supported_extensions = ['PDF', 'DOC', 'DOCX', 'XLSX', 'JPG', 'JPEG', 'PNG', 'WEBP'] file_service = FileProcessingService file_bytes = input_message.file.read() kind = filetype.guess(file_bytes[:550]) if not kind: - if Path(input_message.file.name).suffix[1:].upper() not in supported_extensions: - raise FileExtensionNotSupported(supported_extensions) + if Path(input_message.file.name).suffix[1:].upper() not in self.SUPPORTED_EXTENSIONS: + raise FileExtensionNotSupported(self.SUPPORTED_EXTENSIONS) raise CorruptedFileError raw_file_extension = kind.extension file_extension = file_service.get_file_extension(raw_file_extension, file_bytes) @@ -136,53 +234,6 @@ class Gemini_3_1(SimpleService): {'type': 'image_url', 'image_url': {'url': image_url}}, ] else: - raise FileExtensionNotSupported(supported_extensions) - start_time = time.time() - result = openrouter_run(model_slug, messages, callback_data, 'Gemini 3.1') - process_time = timedelta(seconds=(time.time() - start_time)) - self.handle_invoice( - input_message.content_object.model, - version=version_slug, - input_tokens=result[1], - output_tokens=result[2], - embedding_tokens=embedding_tokens, - ) - msgs = self.save_results(result[0], process_time) - return msgs + raise FileExtensionNotSupported(self.SUPPORTED_EXTENSIONS) - def get_chat_history( - self, message_limit: int = 10, max_character_limit: int = 1500 - ) -> list[dict[str, str | list]]: - if isinstance(self.store, Chat): - air_messages = list( - reversed( - Message.objects.filter( - chats_chats_messages=self.store, is_deleted=False, is_sent=True - ).order_by('-created_at')[1 : message_limit + 1] - ) - ) - elif isinstance(self.store, APIStore): - air_messages = [] - elif isinstance(self.store, Copywrite): - air_messages = list( - reversed( - Message.objects.filter( - copywrite_copywrites_messages=self.store, - is_deleted=False, - is_sent=True, - ).order_by('-created_at')[:message_limit] - ) - ) - else: - air_messages = [] - memory = [] - for msg in air_messages: - content = msg.content or '' - if msg.from_model: - memory.append({'role': 'assistant', 'content': content}) - else: - memory.append({'role': 'user', 'content': content}) - character_length = sum(len(content['content']) for content in memory) - while character_length > max_character_limit: - character_length -= len(memory.pop(0)['content']) - return memory + return messages, embedding_tokens