@@ -0,0 +1,8 @@ +from dataclasses import dataclass + + +@dataclass +class ModelResponse: + content: str + input_tokens: int + output_tokens: int @@ -0,0 +1,85 @@ +import json +import logging + +import httpx + +from backend import settings +from ml_model.utils import count_openrouter_tokens +from poller.models import Proxy + +from .models import ModelResponse + +logger = logging.getLogger(__name__) + + +class OpenrouterAdapter: + BASE_URL = 'https://openrouter.ai/api/v1' + 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 count_openrouter_tokens(model_name, messages, content, bias=bias) @@ -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, @@ -203,6 +204,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,8 @@ -import redis -import tiktoken - +import math from random import randint -from typing import Literal, List, Dict, Any, Tuple +from typing import Any, Dict, List, Literal, Tuple -from django.conf import settings -from redis.commands.search.field import TagField, TextField, VectorField -from redis.commands.search.indexDefinition import IndexDefinition, IndexType +import tiktoken from authentication.models import CustomUserModel from authentication.selectors.account_status_selector import ( @@ -19,12 +15,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(): @@ -39,7 +35,12 @@ def check_account_type( return 'business_account' -def count_openrouter_tokens(model_name: str, messages: List[Dict[str, Any]], output: str) -> Tuple[int, int]: +def count_openrouter_tokens( + model_name: str, + messages: List[Dict[str, Any]], + output: str, + bias: Tuple[float, float] = (1.0, 1.0), +) -> Tuple[int, int]: """A function for count tokens for OpenRouter Neuron Models""" encodings = { 'Qwen': 'cl100k_base', @@ -59,5 +60,7 @@ def count_openrouter_tokens(model_name: str, messages: List[Dict[str, Any]], out else: input_tokens += len(encoding.encode(message['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) -