@@ -1,12 +1,22 @@ +import base64 import time -from _decimal import Decimal +import filetype + +from decimal import Decimal from datetime import timedelta -from typing import Any, Iterator +from io import BytesIO +from typing import Iterator, Any +from PIL import Image -from messages.models import Message -from ml_model.models import ModelCategory, ModelInput, ModelParameter, ModelVersion +from ml_model.models import ModelCategory, ModelVersion, ModelInput from ml_model.services.base import SimpleService -from ml_model.tasks import mistral_run +from ml_model.tasks import openrouter_run + +from tools.chats.models import Chat +from tools.copywrite.models import Copywrite +from tools.public_api.models import APIStore + +from messages.models import Message class Mistral(SimpleService): @@ -18,87 +28,144 @@ class Mistral(SimpleService): title = 'Mistral' description = 'Нейросеть, способная генерировать качественный текст из вашего промпта' category = ModelCategory(title='Чат-боты', slug='chat-bots') + versions = [ - ModelVersion(name='Tiny', slug='mistral-tiny'), - ModelVersion(name='Small', slug='mistral-small'), - ModelVersion(name='Medium', slug='mistral-medium', default=True), + ModelVersion(name='Mistral small', default=True, slug='mistral-small-3.1-24b-instruct'), ] - inputs = [ModelInput(type=ModelInput.TypeChoices.TEXT, required=True)] - parameters = [ - ModelParameter( - name='Лучший процент', - key='top_p', - type=ModelParameter.TypeChoices.FLOATRANGE, - values={'start': 0.0, 'end': 1.0, 'step': 0.1, 'default': 0.5}, - ), - ModelParameter( - name='Температура', - key='temperature', - type=ModelParameter.TypeChoices.FLOATRANGE, - values={'start': 0.0, 'end': 1.0, 'step': 0.1, 'default': 0.5}, - ), + + inputs = [ + ModelInput(type=ModelInput.TypeChoices.TEXT, required=True), + ModelInput(type=ModelInput.TypeChoices.IMAGE) ] - TOKEN_PAYMENT_RULES = { - 'mistral-tiny-input': Decimal(0.165), - 'mistral-tiny-output': Decimal(0.494), - 'mistral-small-input': Decimal(0.706), - 'mistral-small-output': Decimal(2.123), - 'mistral-medium-input': Decimal(2.937), - 'mistral-medium-output': Decimal(8.822), - } + parameters = [] - def __init__(self, store): - super().__init__(store) + TOKENS_COST = { + 'mistral-small-3.1-24b-instruct': { + 'input': Decimal('20'), + 'output': Decimal('60'), + 'input_imgs': Decimal('185.2'), + }, + } - def calculate_price(self, messages: list[Message], input_message: Message) -> Decimal: - price = Decimal( - sum( - [ - self.TOKEN_PAYMENT_RULES[f"{input_message.info['version']}-output"] - / 1000 - * len(msg.content.split(' ')) - for msg in messages - ] - + [ - self.TOKEN_PAYMENT_RULES[f"{input_message.info['version']}-input"] - / 1000 - * len(input_message.content.split(' ')) - ] - ) + def calculate_price( + self, version: str, input_tokens: int, output_tokens: int + ) -> Decimal: + price_map = self.TOKENS_COST[version.split('/')[1]] + price = ( + input_tokens * price_map['input'] / 1_000_000 + + output_tokens * price_map['output'] / 1_000_000 + + price_map['input_imgs'] / 1_000 ) return price.quantize(Decimal('0.1'), rounding='ROUND_UP') - def make_messages(self, r: Iterator[Any], t: timedelta): - msgs = [] - for message in r: - msgs.append( - Message( - content=message['message']['content'], - content_object=self.store, - elapsed_time=t, - ) + def save_results( + self, content: Iterator[Any], 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 save_results(self, messages: list[Message]) -> list[Message]: - return Message.objects.bulk_create(messages) - def make(self, input_message: Message, save: bool = True) -> list[Message]: - info = input_message.info.copy() - version = info.pop('version') - callback_data = dict( - { - 'messages': [{'role': 'user', 'content': input_message.content}], - 'model': version, - **info, - } - ) start_time = time.time() - result = mistral_run(callback_data) + info = input_message.info.copy() + version = f'mistralai/{info.pop('version')}' + callback_data = { + 'provider': { + 'order': ['Parasail'] + }, + **input_message.info + } + messages = self.get_chat_history() + 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 + } + } + ] + result = openrouter_run(version, messages, callback_data, self.title) process_time = timedelta(seconds=(time.time() - start_time)) - msgs = self.make_messages(result['choices'], process_time) - self.handle_invoice(self.neuron_model, messages=msgs, input_message=input_message) + self.handle_invoice( + input_message.content_object.model, + version=version, + input_tokens=result[1], + output_tokens=result[2] + ) + msgs = self.save_results(result[0], process_time) + return msgs + + def __init__(self, store): + super().__init__(store) + + def save_results( + self, content: Iterator[Any], t: timedelta, save: bool = True + ) -> list[Message]: + msgs = [ + Message( + content=content, + content_object=self.store, + elapsed_time=t, + ) + ] if save: - self.save_results(messages=msgs) + return Message.objects.bulk_create(msgs) return msgs + + def get_chat_history(self, message_limit: int = 10, max_character_limit: int = 1500): + 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')[:message_limit] + ) + ) + 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: + memory.pop(0) + character_length = sum(len(content['content']) for content in memory) + return memory @@ -1,5 +1,6 @@ import base64 import json +import httpx # import uuid from io import BytesIO @@ -111,6 +112,37 @@ def replicate_run(callback_url: str, payload: dict[str, Any]): input=payload, ) +@shared_task +def openrouter_run(version: str, messages: list, callback_data: dict, model_name: str): + for proxy in Proxy.objects.all(): + with httpx.Client( + base_url='https://openrouter.ai/api/v1', + headers={'Authorization': f'Bearer {settings.OPENROUTER_API_KEY}'}, + proxy=f'{proxy.protocol}://{proxy.address}', + ) as client: + resp = client.post( + 'chat/completions', + json={ + 'model': version, + 'messages': messages, + **callback_data + }, + ) + if ( + (data := resp.json()) + and data.get('choices') + and ( + content := ','.join( + [choice['message']['content'] for choice in data.get('choices')] + ) + ) + ): + return ( + content, + data['usage']['prompt_tokens'], + data['usage']['completion_tokens'] + ) + raise Exception(f'No answer from {model_name}, please retry later') @shared_task def upscale_run(payload: dict[str, tuple[str, IO]]) -> list[str]: @@ -120,22 +152,6 @@ def upscale_run(payload: dict[str, tuple[str, IO]]) -> list[str]: ).content return json.loads(content) - -@shared_task -def mistral_run(payload: dict[str, Any]): - headers = { - 'Content-Type': 'application/json', - 'Authorization': settings.MISTRAL_API_KEY, - } - return json.loads( - requests.post( - 'https://api.mistral.ai/v1/chat/completions', - json.dumps(payload), - headers=headers, - ).content - ) - - @shared_task def claude_run(payload: dict[str, Any]): headers = {