@@ -0,0 +1,8 @@ +from dataclasses import dataclass + + +@dataclass +class ModelResponse: + content: str + input_tokens: int + output_tokens: int @@ -0,0 +1,121 @@ +import json +import logging +import math +from typing import Any + +import httpx +import tiktoken + +from backend import settings +from poller.models import Proxy + +from .models import ModelResponse + +logger = logging.getLogger(__name__) + + +class OpenrouterAdapter: + BASE_URL = 'https://openrouter.ai/api/v1' + FALLBACK_ENCODINGS: dict[str, str] = { + 'Qwen': 'cl100k_base', + 'Deepseek': 'cl100k_base', + 'Claude': 'r50k_base', + 'Perplexity': 'cl100k_base', + 'Mistral': 'r50k_base', + 'LLaMA': 'cl100k_base', + 'Grok': 'r50k_base', + 'Gemini': 'cl100k_base', + } + FALLBACK_TOKEN_BIAS: dict[str, tuple[float, float]] = { + # Средний калибровочный bias по 15 замерам: + # input: 2814 / 3246 ~= 0.867, output: 14223 / 25767 ~= 0.552 + 'Grok': (0.867, 0.552), + } + + @classmethod + def run_streaming_api( + cls, version: str, messages: list, callback_data: dict, model_name: str + ) -> ModelResponse: + for proxy in Proxy.objects.all(): + with httpx.Client( + base_url=cls.BASE_URL, + headers={'Authorization': f'Bearer {settings.OPENROUTER_API_KEY}'}, + proxy=f'{proxy.protocol}://{proxy.address}', + timeout=600, + ) as client: + with client.stream( + 'POST', + 'chat/completions', + json={ + 'model': version, + 'stream': True, + 'messages': messages, + 'transforms': ['middle-out'], + **callback_data, + }, + ) as resp: + content = '' + # reasoning используем только для фоллбэк-подсчёта токенизатора + # в ответ не кладём, заполняет буфер истории сообщений + reasoning = '' + input_tokens = output_tokens = 0 + for line in resp.iter_lines(): + line = line.strip() + if not line or not line.startswith('data: '): + continue + + data = line[6:] + if data == '[DONE]': + break + + try: + data_obj = json.loads(data) + content += data_obj['choices'][0]['delta'].get('content') or '' + reasoning += data_obj['choices'][0]['delta'].get('reasoning') or '' + if data_obj.get('usage'): + input_tokens = data_obj['usage']['prompt_tokens'] + output_tokens = data_obj['usage']['completion_tokens'] + + except json.JSONDecodeError: + logger.warning( + f'Opernrouter chunk parsing failed for model {model_name}: {data}' + ) + continue + + if not input_tokens or not output_tokens: + logger.error(f'Opernrouter failed get data about tokens for model {model_name}') + input_tokens, output_tokens = cls._fallback_tokenize( + model_name, messages, content + reasoning + ) + + return ModelResponse(content, input_tokens, output_tokens) + + @classmethod + def _fallback_tokenize(cls, model_name: str, messages: list, content: str) -> tuple[int, int]: + bias = cls.FALLBACK_TOKEN_BIAS.get(model_name, (1.0, 1.0)) + return cls.count_tokens_fallback(model_name, messages, content, bias=bias) + + @classmethod + def count_tokens_fallback( + cls, + model_name: str, + messages: list[dict[str, Any]], + output: str, + bias: tuple[float, float] = (1.0, 1.0), + ) -> tuple[int, int]: + """Fallback token counting for OpenRouter responses.""" + encoding = tiktoken.get_encoding(cls.FALLBACK_ENCODINGS[model_name]) + input_tokens = 100 if model_name == 'LLaMA' else 0 + for message in messages: + content = message['content'] + if isinstance(content, list): + input_tokens += len(encoding.encode(content[0]['text'])) + else: + input_tokens += len(encoding.encode(content)) + output_tokens = len(encoding.encode(output)) + + input_bias, output_bias = bias + input_tokens = math.ceil(input_tokens * input_bias) + output_tokens = math.ceil(output_tokens * output_bias) + + return (input_tokens, output_tokens) @@ -3,30 +3,27 @@ import time from datetime import timedelta from decimal import Decimal from io import BytesIO -from typing import Any, Iterator from pathlib import Path +from typing import Any, Iterator import filetype from PIL import Image -from ml_model.services.FileService import FileProcessingService -from ml_model.exceptions import FileExtensionNotSupported, CorruptedFileError + +from messages.models import Message +from ml_model.exceptions import CorruptedFileError, FileExtensionNotSupported +from ml_model.services.base import SimpleService from ml_model.services.EmbeddingService import EmbeddingService +from ml_model.services.FileService import FileProcessingService +from ml_model.tasks import openrouter_run_streaming from payments.exceptions.insufficient_balance import InsufficientBalance from payments.selectors.payment_plan_selector import PaymentPlanSelector - - from poller.models import Proxy - -from messages.models import Message -from ml_model.services.base import SimpleService -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 class Grok(SimpleService): - TOKENS_COST = { 'grok-4.3': { 'input': {'default': Decimal('875'), 'high': Decimal('1750')}, @@ -43,9 +40,7 @@ class Grok(SimpleService): ) -> Decimal: price_map = self.TOKENS_COST[version.split('/')[1]] price = ( - input_tokens - * price_map['input']['default' if input_tokens <= 200_000 else 'high'] - / 1_000_000 + input_tokens * price_map['input']['default' if input_tokens <= 200_000 else 'high'] / 1_000_000 + output_tokens * price_map['output']['default' if output_tokens <= 200_000 else 'high'] / 1_000_000 @@ -135,16 +130,16 @@ class Grok(SimpleService): else: raise FileExtensionNotSupported(self.SUPPORTED_EXTENSIONS) - result = openrouter_run(version, messages, callback_data, 'Grok') + result = openrouter_run_streaming(version, messages, callback_data, 'Grok') process_time = timedelta(seconds=(time.time() - start_time)) self.handle_invoice( input_message.content_object.model, version=version, - input_tokens=result[1], - output_tokens=result[2], + input_tokens=result.input_tokens, + output_tokens=result.output_tokens, embedding_tokens=embedding_tokens, ) - msgs = self.save_results(result[0], process_time) + msgs = self.save_results(result.content, process_time) return msgs def get_chat_history( @@ -0,0 +1 @@ + @@ -1,5 +1,6 @@ from typing import Type +from django.db import transaction from django.db.models.signals import post_save from django.dispatch import receiver @@ -9,4 +10,4 @@ from ml_model.models import ModelSettings, NeuronModel @receiver(post_save, sender=NeuronModel) def create_settings(sender: Type[NeuronModel], instance: NeuronModel, created: bool, **kwargs): if created: - ModelSettings.objects.create(model=instance) + transaction.on_commit(lambda: ModelSettings.objects.get_or_create(model=instance)) @@ -21,6 +21,7 @@ from requests import Response from backend import settings from ml_model.adapters.bytedance_model_ark import BytedanceContentType, BytedanceModelArkAdapter +from ml_model.adapters.openrouter import OpenrouterAdapter from ml_model.exceptions import ( CorruptedFileError, DeploymentDisabled, @@ -35,7 +36,6 @@ from ml_model.exceptions import ( PredictionInterruptedError, RequestBlocked, ) -from ml_model.utils import count_openrouter_tokens from poller.models import Proxy logger = logging.getLogger(__name__) @@ -189,7 +189,7 @@ def openrouter_run(version: str, messages: list, callback_data: dict, model_name error_type = re.sub(r'["\']', '', str(data['choices'][0]['error']['message'])) if error_type == 'Overloaded': logger.warning(f'Model {model_name} overloaded') - input_tokens, output_tokens = count_openrouter_tokens( + input_tokens, output_tokens = OpenrouterAdapter.count_tokens_fallback( model_name, messages, content + reasoning ) else: @@ -203,6 +203,11 @@ def openrouter_run(version: str, messages: list, callback_data: dict, model_name raise Exception(f'No answer from {model_name}, please retry later') +@shared_task +def openrouter_run_streaming(version: str, messages: list, callback_data: dict, model_name: str): + return OpenrouterAdapter.run_streaming_api(version, messages, callback_data, model_name) + + @shared_task def fal_ai_run(model, payload): for proxy in Proxy.objects.all(): @@ -1,12 +1,5 @@ -import redis -import tiktoken - from random import randint -from typing import Literal, List, Dict, Any, Tuple - -from django.conf import settings -from redis.commands.search.field import TagField, TextField, VectorField -from redis.commands.search.indexDefinition import IndexDefinition, IndexType +from typing import Literal from authentication.models import CustomUserModel from authentication.selectors.account_status_selector import ( @@ -19,12 +12,12 @@ from authentication.selectors.business_account_selector import ( def random_with_N_digits(n): range_start = 10 ** (n - 1) - range_end = (10 ** n) - 1 + range_end = (10**n) - 1 return randint(range_start, range_end) def check_account_type( - user: CustomUserModel, + user: CustomUserModel, ) -> Literal['business_host'] | Literal['business_account'] | Literal['regular'] | Literal['business_admin']: status = AccountStatusSelector(user) if status.is_business_host(): @@ -37,27 +30,3 @@ def check_account_type( return 'business_admin' return 'business_account' - - -def count_openrouter_tokens(model_name: str, messages: List[Dict[str, Any]], output: str) -> Tuple[int, int]: - """A function for count tokens for OpenRouter Neuron Models""" - encodings = { - 'Qwen': 'cl100k_base', - 'Deepseek': 'cl100k_base', - 'Claude': 'r50k_base', - 'Perplexity': 'cl100k_base', - 'Mistral': 'r50k_base', - 'LLaMA': 'cl100k_base', - 'Grok': 'r50k_base', - 'Gemini': 'cl100k_base', - } - encoding = tiktoken.get_encoding(encodings[model_name]) - input_tokens = 100 if model_name == 'LLaMA' else 0 - for message in messages: - if isinstance(message['content'], list): - input_tokens += len(encoding.encode(message['content'][0]['text'])) - else: - input_tokens += len(encoding.encode(message['content'])) - output_tokens = len(encoding.encode(output)) - return (input_tokens, output_tokens) -